Compare commits
9
Commits
9633ce99d2
..
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2421cdc273 | ||
|
|
f35e3ddfe8 | ||
|
|
0dc26041dc | ||
|
|
c80753c56e | ||
|
|
26f4abef7f | ||
|
|
f35cbd01cf | ||
|
|
70a8ea1b92 | ||
|
|
432097ee3f | ||
|
|
5d6f6ffaf9 |
@@ -61,6 +61,10 @@ linters:
|
|||||||
desc: >-
|
desc: >-
|
||||||
Test-support code belongs in test files and in packages whose
|
Test-support code belongs in test files and in packages whose
|
||||||
directory name ends in test, not in the shipped binary.
|
directory name ends in test, not in the shipped binary.
|
||||||
|
- pkg: sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest
|
||||||
|
desc: >-
|
||||||
|
Test-support code belongs in test files and in packages whose
|
||||||
|
directory name ends in test, not in the shipped binary.
|
||||||
# Only decisions already recorded in the Go package defaults are
|
# Only decisions already recorded in the Go package defaults are
|
||||||
# listed here. Every entry matches the module path exactly.
|
# listed here. Every entry matches the module path exactly.
|
||||||
gomodguard_v2:
|
gomodguard_v2:
|
||||||
|
|||||||
+25
-1
@@ -29,6 +29,12 @@ RUN go mod download
|
|||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
|
# go.mod and go.sum must be as `go mod tidy` writes them, which is what
|
||||||
|
# `make tidy` does. Checked before the tests, which a missing go.sum line
|
||||||
|
# fails with a message that does not name `make tidy`.
|
||||||
|
RUN go mod tidy -diff || \
|
||||||
|
{ echo "go.mod or go.sum is not tidy: run make tidy" >&2; exit 1; }
|
||||||
|
|
||||||
# Go's build cache is kept on a tmpfs, out of the image: nothing uses it
|
# Go's build cache is kept on a tmpfs, out of the image: nothing uses it
|
||||||
# after this step, and writing it into the image takes seconds.
|
# after this step, and writing it into the image takes seconds.
|
||||||
RUN --mount=type=tmpfs,target=/root/.cache/go-build \
|
RUN --mount=type=tmpfs,target=/root/.cache/go-build \
|
||||||
@@ -36,7 +42,25 @@ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
|
|||||||
{ echo "--- Rerunning with -v for details ---"; \
|
{ echo "--- Rerunning with -v for details ---"; \
|
||||||
go test -timeout 90s -race -v ./...; exit 1; }
|
go test -timeout 90s -race -v ./...; exit 1; }
|
||||||
|
|
||||||
# Build stage. Nothing is wanted from the two phases above; the copies
|
# Tidy stage: `go mod tidy` in the test phase's Go, so that the files it
|
||||||
|
# writes pass the test phase's check. Nothing else depends on it, so only
|
||||||
|
# script/tidy, which names the stage after it, builds it.
|
||||||
|
#
|
||||||
|
# golang 1.27.1-trixie, 2026-09-19
|
||||||
|
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS tidy
|
||||||
|
|
||||||
|
WORKDIR /src
|
||||||
|
|
||||||
|
COPY . .
|
||||||
|
|
||||||
|
RUN go mod tidy
|
||||||
|
|
||||||
|
# go.mod and go.sum alone, which script/tidy writes into the working tree.
|
||||||
|
FROM scratch AS tidy-files
|
||||||
|
|
||||||
|
COPY --from=tidy /src/go.mod /src/go.sum /
|
||||||
|
|
||||||
|
# Build stage. Nothing is wanted from the lint and test phases; the copies
|
||||||
# are what make BuildKit build them first, so the image, which needs this
|
# are what make BuildKit build them first, so the image, which needs this
|
||||||
# stage, cannot be produced unless lint and test passed.
|
# stage, cannot be produced unless lint and test passed.
|
||||||
#
|
#
|
||||||
|
|||||||
@@ -1,9 +1,10 @@
|
|||||||
.PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \
|
.PHONY: bootstrap setup test lint fmt fmt-check tidy check docker hooks build \
|
||||||
example-app
|
run example-app
|
||||||
|
|
||||||
# Makefile targets are thin shims; the implementations live in script/
|
# Makefile targets are thin shims; the implementations live in script/
|
||||||
# per the scripts-to-rule-them-all pattern (see the Entrypoints section
|
# per the scripts-to-rule-them-all pattern (see the Entrypoints section
|
||||||
# of README.md). build and run are for working on the code by hand;
|
# of README.md). tidy writes go.mod and go.sum as `go mod tidy` does,
|
||||||
|
# which test checks. build and run are for working on the code by hand;
|
||||||
# example-app checks the image with an app built on it.
|
# example-app checks the image with an app built on it.
|
||||||
|
|
||||||
bootstrap:
|
bootstrap:
|
||||||
@@ -24,6 +25,9 @@ fmt:
|
|||||||
fmt-check:
|
fmt-check:
|
||||||
@script/fmt-check
|
@script/fmt-check
|
||||||
|
|
||||||
|
tidy:
|
||||||
|
@script/tidy
|
||||||
|
|
||||||
check:
|
check:
|
||||||
@script/check
|
@script/check
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ go 1.26.0
|
|||||||
require (
|
require (
|
||||||
github.com/fsnotify/fsnotify v1.10.1
|
github.com/fsnotify/fsnotify v1.10.1
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7
|
github.com/hashicorp/golang-lru/v2 v2.0.7
|
||||||
|
github.com/maxmind/mmdbwriter v1.2.0
|
||||||
|
github.com/oschwald/maxminddb-golang/v2 v2.7.0
|
||||||
github.com/prometheus/client_golang v1.24.1
|
github.com/prometheus/client_golang v1.24.1
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -16,6 +18,7 @@ require (
|
|||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
github.com/prometheus/common v0.70.1 // indirect
|
github.com/prometheus/common v0.70.1 // indirect
|
||||||
github.com/prometheus/procfs v0.21.1 // indirect
|
github.com/prometheus/procfs v0.21.1 // indirect
|
||||||
golang.org/x/sys v0.47.0 // indirect
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
|
||||||
|
golang.org/x/sys v0.48.0 // indirect
|
||||||
google.golang.org/protobuf v1.36.11 // indirect
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,8 +2,6 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
|||||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
|
||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
|
||||||
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||||
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
@@ -14,10 +12,12 @@ github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJn
|
|||||||
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||||
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||||
|
github.com/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM=
|
||||||
|
github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U=
|
||||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g=
|
||||||
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
||||||
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
||||||
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
@@ -26,15 +26,17 @@ github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi
|
|||||||
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
||||||
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
||||||
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
||||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
||||||
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
|
||||||
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
||||||
|
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||||
|
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
|
||||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
|
||||||
|
|||||||
@@ -0,0 +1,884 @@
|
|||||||
|
// Package alerts sends alerts on bans, on a source that fails and on a
|
||||||
|
// file with an error to each destination set: to the webhook
|
||||||
|
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
|
||||||
|
// webhook schema" section of SPEC.md describes, to the Slack incoming
|
||||||
|
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
|
||||||
|
// ntfy topic SWWAF_ALERT_NTFY_URL names. A repeat within
|
||||||
|
// SWWAF_ALERT_COOLDOWN is held back, and so is an alert past
|
||||||
|
// SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary. The others wait in a
|
||||||
|
// bounded queue of each destination's own, so that a destination that is
|
||||||
|
// slow or unreachable holds up neither the others nor any request. The
|
||||||
|
// state is written to alerts.json and read from it by the state package.
|
||||||
|
// Nothing logged names a destination's URL, whose path or query can carry
|
||||||
|
// a secret.
|
||||||
|
package alerts
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"cmp"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"maps"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"net/url"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The events an alert is for, as SWWAF_ALERT_EVENTS names them.
|
||||||
|
const (
|
||||||
|
// EventBan is a ban smallwebwaf made.
|
||||||
|
EventBan = "ban"
|
||||||
|
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
|
||||||
|
// made permanent.
|
||||||
|
EventPermanentBan = "permanent_ban"
|
||||||
|
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
|
||||||
|
// Core Rule Set, the anomaly thresholds and the reputation sources;
|
||||||
|
// nothing raises them yet.
|
||||||
|
EventWAFBlock = "waf_block"
|
||||||
|
EventAnomaly = "anomaly"
|
||||||
|
EventReputationHit = "reputation_hit"
|
||||||
|
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
|
||||||
|
EventSourceFailure = "source_failure"
|
||||||
|
// EventFileError is a rule file or state file edited while smallwebwaf
|
||||||
|
// runs that does not parse, a replacement of the lookup database that
|
||||||
|
// cannot be read, or a state file that cannot be written.
|
||||||
|
EventFileError = "file_error"
|
||||||
|
// EventSummary is the summary of the alerts an hour held back past
|
||||||
|
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
|
||||||
|
EventSummary = "summary"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Events returns every event SWWAF_ALERT_EVENTS can name, which is its
|
||||||
|
// default.
|
||||||
|
func Events() []string {
|
||||||
|
return []string{
|
||||||
|
EventBan, EventPermanentBan, EventWAFBlock, EventAnomaly,
|
||||||
|
EventReputationHit, EventSourceFailure, EventFileError,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The destinations alerts are sent to, as the metrics and alerts.json
|
||||||
|
// name them.
|
||||||
|
const (
|
||||||
|
// DestinationWebhook is the webhook SWWAF_ALERT_WEBHOOK_URL names.
|
||||||
|
DestinationWebhook = "webhook"
|
||||||
|
// DestinationSlack is the Slack incoming webhook
|
||||||
|
// SWWAF_ALERT_SLACK_WEBHOOK_URL names.
|
||||||
|
DestinationSlack = "slack"
|
||||||
|
// DestinationNtfy is the ntfy topic SWWAF_ALERT_NTFY_URL names.
|
||||||
|
DestinationNtfy = "ntfy"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Destinations returns every destination alerts can be sent to.
|
||||||
|
func Destinations() []string {
|
||||||
|
return []string{DestinationWebhook, DestinationSlack, DestinationNtfy}
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
// queueSize is the most alerts that wait to be sent to a destination.
|
||||||
|
// Past it, the oldest is dropped.
|
||||||
|
queueSize = 1000
|
||||||
|
// sendTimeout bounds one request to a destination.
|
||||||
|
sendTimeout = 10 * time.Second
|
||||||
|
// After a request to a destination fails, the alert is sent again a
|
||||||
|
// second later, and retryDelayFactor times as long after each further
|
||||||
|
// failure in a row, up to a minute.
|
||||||
|
firstRetryDelay = time.Second
|
||||||
|
retryDelayFactor = 2
|
||||||
|
maxRetryDelay = time.Minute
|
||||||
|
// maxAnswerBytes is the most of a destination's answer that is read.
|
||||||
|
maxAnswerBytes = 64 << 10
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errStatus = errors.New("the destination answered")
|
||||||
|
// errRefused is a 4xx answer other than 408 and 429: the destination
|
||||||
|
// refuses the alert itself, and would refuse it again.
|
||||||
|
errRefused = errors.New("the destination refused the alert, answering")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Params are what New needs. With none of WebhookURL, SlackURL and
|
||||||
|
// NtfyURL set, no alert is sent.
|
||||||
|
type Params struct {
|
||||||
|
// WebhookURL is where each alert is posted as JSON
|
||||||
|
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. WebhookHeaders
|
||||||
|
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
|
||||||
|
WebhookURL *url.URL
|
||||||
|
WebhookHeaders http.Header
|
||||||
|
// SlackURL is the Slack incoming webhook each alert is posted to as a
|
||||||
|
// message (SWWAF_ALERT_SLACK_WEBHOOK_URL), nil while it is unset.
|
||||||
|
SlackURL *url.URL
|
||||||
|
// NtfyURL is the ntfy topic each alert is published to
|
||||||
|
// (SWWAF_ALERT_NTFY_URL), nil while it is unset. NtfyToken, unless
|
||||||
|
// empty, is sent with each as a bearer token (SWWAF_ALERT_NTFY_TOKEN).
|
||||||
|
NtfyURL *url.URL
|
||||||
|
NtfyToken string
|
||||||
|
// Events are the events alerts are sent for (SWWAF_ALERT_EVENTS).
|
||||||
|
Events []string
|
||||||
|
// Cooldown is how long a repeat of an alert is held back
|
||||||
|
// (SWWAF_ALERT_COOLDOWN), 0 for no time. MaxPerHour is the most alerts
|
||||||
|
// sent in an hour (SWWAF_ALERT_MAX_PER_HOUR), 0 for no limit.
|
||||||
|
Cooldown time.Duration
|
||||||
|
MaxPerHour int
|
||||||
|
// Instance is SWWAF_INSTANCE_NAME, which every alert gives.
|
||||||
|
Instance string
|
||||||
|
// Now tells the time of an alert, normally time.Now in UTC.
|
||||||
|
Now func() time.Time
|
||||||
|
// ProcessLog receives the requests to a destination that fail.
|
||||||
|
ProcessLog *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// Alert is one alert, as the webhook is sent it and alerts.json holds it,
|
||||||
|
// with the fields of the "Alert webhook schema" section of SPEC.md. ASN,
|
||||||
|
// ASName and Country are, for a ban, the client's as the ban's notes give
|
||||||
|
// them.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // SPEC.md's alert webhook schema names its fields in snake_case
|
||||||
|
type Alert struct {
|
||||||
|
Instance string `json:"instance"`
|
||||||
|
Time time.Time `json:"time"`
|
||||||
|
Event string `json:"event"`
|
||||||
|
Client netip.Addr `json:"client"`
|
||||||
|
Netblock netip.Prefix `json:"netblock"`
|
||||||
|
ASN string `json:"asn"`
|
||||||
|
ASName string `json:"as_name"`
|
||||||
|
Country string `json:"country"`
|
||||||
|
// Reason is a short sentence, and Detail what is particular to the
|
||||||
|
// event: for a file_error, its "file", and for a source_failure, its
|
||||||
|
// "source", which the cooldown tells repeats by.
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
Detail map[string]any `json:"detail"`
|
||||||
|
// SuppressedRepeats is how many repeats of the alert the cooldown
|
||||||
|
// held back since the last one let through.
|
||||||
|
SuppressedRepeats int `json:"suppressed_repeats"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown is, for an event on a netblock, or about a file or a source,
|
||||||
|
// when the last alert let through was raised, and how many repeats the
|
||||||
|
// cooldown has held back since, as alerts.json holds it.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
|
type Cooldown struct {
|
||||||
|
Event string `json:"event"`
|
||||||
|
Netblock netip.Prefix `json:"netblock"`
|
||||||
|
File string `json:"file,omitempty"`
|
||||||
|
Source string `json:"source,omitempty"`
|
||||||
|
Sent time.Time `json:"sent"`
|
||||||
|
SuppressedRepeats int `json:"suppressed_repeats"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hour is the hour under way, by the clock, as alerts.json holds it: when
|
||||||
|
// it started, how many alerts were let through in it, and how many were
|
||||||
|
// held back in it past MaxPerHour, by event, for its summary.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
|
type Hour struct {
|
||||||
|
Start time.Time `json:"start"`
|
||||||
|
Sent int `json:"sent"`
|
||||||
|
HeldBack map[string]int `json:"held_back"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// State is what alerts.json holds: the cooldowns, the hour under way, and
|
||||||
|
// for each destination set, the alerts waiting to be sent to it, oldest
|
||||||
|
// first.
|
||||||
|
type State struct {
|
||||||
|
Cooldowns []Cooldown `json:"cooldowns"`
|
||||||
|
Hour Hour `json:"hour"`
|
||||||
|
Waiting map[string][]Alert `json:"waiting"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Counts are, for a destination, how many alerts it took, how many
|
||||||
|
// requests to it failed, and how many alerts were dropped from its full
|
||||||
|
// queue or given up as it refused them.
|
||||||
|
type Counts struct {
|
||||||
|
Sent, Failed, Dropped int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// Queue takes the alerts raised, holds back those it must, and sends the
|
||||||
|
// others to each destination set, from a queue of the destination's own.
|
||||||
|
// It is safe for concurrent use.
|
||||||
|
type Queue struct {
|
||||||
|
params Params
|
||||||
|
// destinations are the destinations set, in the order of
|
||||||
|
// Destinations.
|
||||||
|
destinations []*destination
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
// cooldowns are the alerts last let through, by event and netblock,
|
||||||
|
// file or source.
|
||||||
|
cooldowns map[cooldownKey]*Cooldown
|
||||||
|
hour Hour
|
||||||
|
|
||||||
|
suppressed atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// destination is a destination set, with the alerts waiting to be sent
|
||||||
|
// to it. Its mu is taken after the Queue's, never before.
|
||||||
|
type destination struct {
|
||||||
|
// name is how the metrics and alerts.json name the destination, and
|
||||||
|
// setting the setting that is its URL, which the log names in place
|
||||||
|
// of the URL.
|
||||||
|
name string
|
||||||
|
setting string
|
||||||
|
url *url.URL
|
||||||
|
// message returns the body an alert is posted with, and the headers
|
||||||
|
// sent with it.
|
||||||
|
message func(alert *Alert) ([]byte, http.Header, error)
|
||||||
|
// httpClient follows no redirect: a redirect is a failure.
|
||||||
|
httpClient *http.Client
|
||||||
|
processLog *slog.Logger
|
||||||
|
// queued receives a value when an alert joins the queue, unless one
|
||||||
|
// waits already, so that run looks at the queue again.
|
||||||
|
queued chan struct{}
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
// waiting are the alerts waiting to be sent, oldest first.
|
||||||
|
waiting []*Alert
|
||||||
|
|
||||||
|
sent, failed, dropped atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// cooldownKey is what makes an alert a repeat of another: the same event
|
||||||
|
// on the same netblock, and about the same file or source, as its detail
|
||||||
|
// names them. Each is empty for an alert without one.
|
||||||
|
type cooldownKey struct {
|
||||||
|
event string
|
||||||
|
netblock netip.Prefix
|
||||||
|
file string
|
||||||
|
source string
|
||||||
|
}
|
||||||
|
|
||||||
|
// cooldownKeyOf returns what makes another alert a repeat of alert.
|
||||||
|
func cooldownKeyOf(alert *Alert) cooldownKey {
|
||||||
|
file, _ := alert.Detail["file"].(string)
|
||||||
|
source, _ := alert.Detail["source"].(string)
|
||||||
|
|
||||||
|
return cooldownKey{alert.Event, alert.Netblock, file, source}
|
||||||
|
}
|
||||||
|
|
||||||
|
// New returns a Queue with no alert yet.
|
||||||
|
func New(params Params) *Queue {
|
||||||
|
q := &Queue{
|
||||||
|
params: params,
|
||||||
|
cooldowns: map[cooldownKey]*Cooldown{},
|
||||||
|
hour: Hour{HeldBack: map[string]int{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
if params.WebhookURL != nil {
|
||||||
|
q.addDestination(DestinationWebhook, "SWWAF_ALERT_WEBHOOK_URL",
|
||||||
|
params.WebhookURL, q.webhookMessage)
|
||||||
|
}
|
||||||
|
|
||||||
|
if params.SlackURL != nil {
|
||||||
|
q.addDestination(DestinationSlack, "SWWAF_ALERT_SLACK_WEBHOOK_URL",
|
||||||
|
params.SlackURL, slackMessage)
|
||||||
|
}
|
||||||
|
|
||||||
|
if params.NtfyURL != nil {
|
||||||
|
q.addDestination(DestinationNtfy, "SWWAF_ALERT_NTFY_URL",
|
||||||
|
params.NtfyURL, q.ntfyMessage)
|
||||||
|
}
|
||||||
|
|
||||||
|
return q
|
||||||
|
}
|
||||||
|
|
||||||
|
// Raise sends alert, which names its event and what is particular to it,
|
||||||
|
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event
|
||||||
|
// out. It gives alert the instance and the time. An alert that repeats
|
||||||
|
// the last one let through less than Cooldown before is held back and
|
||||||
|
// counted, and the next one let through gives that count. Past MaxPerHour
|
||||||
|
// alerts let through in the hour under way, by the clock, an alert is
|
||||||
|
// held back for that hour's summary instead, which is sent once the hour
|
||||||
|
// has ended; it starts no cooldown, and the repeats held back before it
|
||||||
|
// are given by the next alert let through. Raise never waits: an alert
|
||||||
|
// let through joins the queue of each destination, from which Run sends
|
||||||
|
// it, and with queueSize alerts waiting for a destination, the oldest is
|
||||||
|
// dropped.
|
||||||
|
func (q *Queue) Raise(alert Alert) {
|
||||||
|
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
|
now := q.params.Now()
|
||||||
|
alert.Instance = q.params.Instance
|
||||||
|
alert.Time = now
|
||||||
|
|
||||||
|
if q.repeat(&alert, now) {
|
||||||
|
q.suppressed.Add(1)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
q.endHour(now)
|
||||||
|
|
||||||
|
if q.params.MaxPerHour > 0 && q.hour.Sent >= q.params.MaxPerHour {
|
||||||
|
q.hour.HeldBack[alert.Event]++
|
||||||
|
q.suppressed.Add(1)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
q.startCooldown(&alert, now)
|
||||||
|
q.hour.Sent++
|
||||||
|
q.queue(&alert)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WouldSend reports whether Raise would let an alert for event on
|
||||||
|
// netblock through now: a destination is set, SWWAF_ALERT_EVENTS chooses
|
||||||
|
// event, no alert for event on netblock was let through less than
|
||||||
|
// Cooldown before, and fewer than MaxPerHour alerts have been let through
|
||||||
|
// in the hour under way. Unlike Raise, it counts nothing.
|
||||||
|
func (q *Queue) WouldSend(event string, netblock netip.Prefix) bool {
|
||||||
|
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, event) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
|
now := q.params.Now()
|
||||||
|
|
||||||
|
last, found := q.cooldowns[cooldownKey{event: event, netblock: netblock}]
|
||||||
|
if q.params.Cooldown > 0 && found && now.Sub(last.Sent) < q.params.Cooldown {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
q.endHour(now)
|
||||||
|
|
||||||
|
return q.params.MaxPerHour == 0 || q.hour.Sent < q.params.MaxPerHour
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run sends the alerts waiting to each destination, from its own queue,
|
||||||
|
// as destination.run does, until ctx is done. It also ends each hour as
|
||||||
|
// Raise does, so that the hour's summary is sent as it ends. With no
|
||||||
|
// destination set, it returns at once.
|
||||||
|
func (q *Queue) Run(ctx context.Context) {
|
||||||
|
if len(q.destinations) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var sending sync.WaitGroup
|
||||||
|
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
sending.Go(func() { d.run(ctx) })
|
||||||
|
}
|
||||||
|
|
||||||
|
for {
|
||||||
|
q.mu.Lock()
|
||||||
|
untilHourEnds := q.hour.Start.Add(time.Hour).Sub(q.params.Now())
|
||||||
|
q.mu.Unlock()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
sending.Wait()
|
||||||
|
|
||||||
|
return
|
||||||
|
case <-time.After(untilHourEnds):
|
||||||
|
q.mu.Lock()
|
||||||
|
q.endHour(q.params.Now())
|
||||||
|
q.mu.Unlock()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Counts returns the counts of the destination name, all 0 for one not
|
||||||
|
// set.
|
||||||
|
func (q *Queue) Counts(name string) Counts {
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
if d.name == name {
|
||||||
|
return Counts{
|
||||||
|
Sent: d.sent.Load(), Failed: d.failed.Load(), Dropped: d.dropped.Load(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return Counts{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// DestinationsSet returns the destinations set, in the order of
|
||||||
|
// Destinations.
|
||||||
|
func (q *Queue) DestinationsSet() []string {
|
||||||
|
names := make([]string, 0, len(q.destinations))
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
names = append(names, d.name)
|
||||||
|
}
|
||||||
|
|
||||||
|
return names
|
||||||
|
}
|
||||||
|
|
||||||
|
// Suppressed is how many alerts were held back: by the cooldown, and past
|
||||||
|
// MaxPerHour. No destination is sent such an alert.
|
||||||
|
func (q *Queue) Suppressed() int64 {
|
||||||
|
return q.suppressed.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot returns the queue's state, as alerts.json holds it, with the
|
||||||
|
// cooldowns sorted by netblock, then by event, file and source.
|
||||||
|
func (q *Queue) Snapshot() State {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
|
state := State{
|
||||||
|
Cooldowns: make([]Cooldown, 0, len(q.cooldowns)),
|
||||||
|
Hour: q.hour,
|
||||||
|
Waiting: map[string][]Alert{},
|
||||||
|
}
|
||||||
|
state.Hour.HeldBack = maps.Clone(q.hour.HeldBack)
|
||||||
|
|
||||||
|
for _, cooldown := range q.cooldowns {
|
||||||
|
state.Cooldowns = append(state.Cooldowns, *cooldown)
|
||||||
|
}
|
||||||
|
|
||||||
|
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
|
||||||
|
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
|
||||||
|
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
state.Waiting[d.name] = d.snapshot()
|
||||||
|
}
|
||||||
|
|
||||||
|
return state
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load puts state, read from alerts.json, in place of the queue's state.
|
||||||
|
// Each cooldown's netblock is masked to its length, so that
|
||||||
|
// 203.0.113.9/24 is 203.0.113.0/24. The alerts waiting for a destination
|
||||||
|
// that is not set are dropped, and so are the oldest past queueSize
|
||||||
|
// alerts waiting for one that is.
|
||||||
|
func (q *Queue) Load(state State) {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
|
||||||
|
q.cooldowns = map[cooldownKey]*Cooldown{}
|
||||||
|
|
||||||
|
for _, cooldown := range state.Cooldowns {
|
||||||
|
cooldown.Netblock = cooldown.Netblock.Masked()
|
||||||
|
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
|
||||||
|
q.cooldowns[key] = &cooldown
|
||||||
|
}
|
||||||
|
|
||||||
|
q.hour = state.Hour
|
||||||
|
q.hour.HeldBack = maps.Clone(state.Hour.HeldBack)
|
||||||
|
|
||||||
|
if q.hour.HeldBack == nil {
|
||||||
|
q.hour.HeldBack = map[string]int{}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
d.load(state.Waiting[d.name])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// addDestination adds a destination: its name, the setting that gives
|
||||||
|
// its URL, that URL, target, and message, which makes the messages sent
|
||||||
|
// to it.
|
||||||
|
func (q *Queue) addDestination(
|
||||||
|
name, setting string, target *url.URL,
|
||||||
|
message func(alert *Alert) ([]byte, http.Header, error),
|
||||||
|
) {
|
||||||
|
q.destinations = append(q.destinations, &destination{
|
||||||
|
name: name,
|
||||||
|
setting: setting,
|
||||||
|
url: target,
|
||||||
|
message: message,
|
||||||
|
httpClient: &http.Client{
|
||||||
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||||
|
return http.ErrUseLastResponse
|
||||||
|
},
|
||||||
|
},
|
||||||
|
processLog: q.params.ProcessLog,
|
||||||
|
queued: make(chan struct{}, 1),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// repeat reports whether alert, raised at now, repeats the last one let
|
||||||
|
// through less than Cooldown before, and counts it if it does.
|
||||||
|
func (q *Queue) repeat(alert *Alert, now time.Time) bool {
|
||||||
|
if q.params.Cooldown == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
last, found := q.cooldowns[cooldownKeyOf(alert)]
|
||||||
|
if !found || now.Sub(last.Sent) >= q.params.Cooldown {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
last.SuppressedRepeats++
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// startCooldown gives alert, let through at now, the count of the repeats
|
||||||
|
// held back since the last one let through, and notes alert as the last
|
||||||
|
// one let through.
|
||||||
|
func (q *Queue) startCooldown(alert *Alert, now time.Time) {
|
||||||
|
if q.params.Cooldown == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
key := cooldownKeyOf(alert)
|
||||||
|
|
||||||
|
last, found := q.cooldowns[key]
|
||||||
|
if found {
|
||||||
|
alert.SuppressedRepeats = last.SuppressedRepeats
|
||||||
|
}
|
||||||
|
|
||||||
|
q.cooldowns[key] = &Cooldown{
|
||||||
|
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
|
||||||
|
Sent: now,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// endHour ends the hour under way, if now is past it: it queues that
|
||||||
|
// hour's summary when alerts were held back in it past MaxPerHour, and
|
||||||
|
// forgets the cooldowns that have run out with no repeat held back, which
|
||||||
|
// no alert needs any more.
|
||||||
|
func (q *Queue) endHour(now time.Time) {
|
||||||
|
start := now.Truncate(time.Hour)
|
||||||
|
if !start.After(q.hour.Start) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
heldBack := 0
|
||||||
|
for _, count := range q.hour.HeldBack {
|
||||||
|
heldBack += count
|
||||||
|
}
|
||||||
|
|
||||||
|
if heldBack > 0 {
|
||||||
|
q.queue(&Alert{
|
||||||
|
Instance: q.params.Instance,
|
||||||
|
Time: now,
|
||||||
|
Event: EventSummary,
|
||||||
|
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
|
||||||
|
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
|
||||||
|
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
|
||||||
|
Detail: map[string]any{
|
||||||
|
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
|
||||||
|
|
||||||
|
for key, cooldown := range q.cooldowns {
|
||||||
|
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
|
||||||
|
delete(q.cooldowns, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// queue adds alert to the alerts waiting for each destination.
|
||||||
|
func (q *Queue) queue(alert *Alert) {
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
d.add(alert)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// run sends the alerts waiting, oldest first, until ctx is done. An alert
|
||||||
|
// stays in the queue until the destination answers it with a 2xx status,
|
||||||
|
// or refuses it with a 4xx status other than 408 and 429: a refused alert
|
||||||
|
// is logged, counted as dropped, and given up, so that the next is sent.
|
||||||
|
// Any other request that fails is logged, and the alert sent again
|
||||||
|
// firstRetryDelay later, retryDelayFactor times as long after each
|
||||||
|
// further failure in a row, up to maxRetryDelay.
|
||||||
|
func (d *destination) run(ctx context.Context) {
|
||||||
|
var (
|
||||||
|
retryDelay time.Duration
|
||||||
|
retryAt time.Time
|
||||||
|
)
|
||||||
|
|
||||||
|
for {
|
||||||
|
alert := d.oldest()
|
||||||
|
|
||||||
|
var due <-chan time.Time // nil while no alert waits
|
||||||
|
if alert != nil {
|
||||||
|
due = time.After(time.Until(retryAt))
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-d.queued:
|
||||||
|
case <-due:
|
||||||
|
err := d.send(ctx, alert)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case err == nil:
|
||||||
|
d.remove(alert)
|
||||||
|
d.sent.Add(1)
|
||||||
|
|
||||||
|
retryDelay = 0
|
||||||
|
retryAt = time.Time{}
|
||||||
|
case errors.Is(err, errRefused):
|
||||||
|
d.remove(alert)
|
||||||
|
d.failed.Add(1)
|
||||||
|
d.dropped.Add(1)
|
||||||
|
|
||||||
|
retryDelay = 0
|
||||||
|
retryAt = time.Time{}
|
||||||
|
|
||||||
|
d.processLog.Warn("gave up an alert "+d.setting+" refused",
|
||||||
|
"event", alert.Event, "error", err.Error())
|
||||||
|
case ctx.Err() == nil: // not cut off as smallwebwaf stops
|
||||||
|
d.failed.Add(1)
|
||||||
|
|
||||||
|
retryDelay = min(max(retryDelayFactor*retryDelay, firstRetryDelay),
|
||||||
|
maxRetryDelay)
|
||||||
|
retryAt = time.Now().Add(retryDelay)
|
||||||
|
|
||||||
|
d.processLog.Warn("sending an alert to "+d.setting+" failed",
|
||||||
|
"error", err.Error(), "sending_again_in", retryDelay.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// add adds alert to the alerts waiting, first dropping the oldest while
|
||||||
|
// queueSize wait, and has run look at the queue again.
|
||||||
|
func (d *destination) add(alert *Alert) {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
|
||||||
|
if len(d.waiting) == queueSize {
|
||||||
|
d.waiting = slices.Delete(d.waiting, 0, 1)
|
||||||
|
d.dropped.Add(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
d.waiting = append(d.waiting, alert)
|
||||||
|
|
||||||
|
select {
|
||||||
|
case d.queued <- struct{}{}:
|
||||||
|
default: // a value waits already
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// load puts waiting, read from alerts.json, in place of the alerts
|
||||||
|
// waiting, as add adds them.
|
||||||
|
func (d *destination) load(waiting []Alert) {
|
||||||
|
d.mu.Lock()
|
||||||
|
d.waiting = nil
|
||||||
|
d.mu.Unlock()
|
||||||
|
|
||||||
|
for _, alert := range waiting {
|
||||||
|
d.add(&alert)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// snapshot returns the alerts waiting, oldest first.
|
||||||
|
func (d *destination) snapshot() []Alert {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
|
||||||
|
waiting := make([]Alert, 0, len(d.waiting))
|
||||||
|
for _, alert := range d.waiting {
|
||||||
|
waiting = append(waiting, *alert)
|
||||||
|
}
|
||||||
|
|
||||||
|
return waiting
|
||||||
|
}
|
||||||
|
|
||||||
|
// oldest returns the oldest alert waiting, nil when none waits.
|
||||||
|
func (d *destination) oldest() *Alert {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
|
||||||
|
if len(d.waiting) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return d.waiting[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
// remove takes alert, which run has sent or given up, out of the queue,
|
||||||
|
// unless it has been dropped from it, or load has replaced the queue,
|
||||||
|
// since run took it. Only the oldest alert is ever dropped, so alert is
|
||||||
|
// the oldest if it is there at all.
|
||||||
|
func (d *destination) remove(alert *Alert) {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
|
||||||
|
if len(d.waiting) > 0 && d.waiting[0] == alert {
|
||||||
|
d.waiting = slices.Delete(d.waiting, 0, 1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// send posts alert to the destination, as message makes it, and returns
|
||||||
|
// an error unless the destination answers with a 2xx status: one that
|
||||||
|
// wraps errRefused for a 4xx status other than 408 and 429. No error
|
||||||
|
// names the destination's URL, whose path or query can carry a secret.
|
||||||
|
func (d *destination) send(ctx context.Context, alert *Alert) error {
|
||||||
|
body, header, err := d.message(alert)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(ctx, sendTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, d.url.String(),
|
||||||
|
bytes.NewReader(body))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("make the request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
maps.Copy(req.Header, header)
|
||||||
|
|
||||||
|
res, err := d.httpClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
// The client's error names the URL: only what went wrong is kept.
|
||||||
|
if urlErr, ok := errors.AsType[*url.Error](err); ok {
|
||||||
|
return urlErr.Err
|
||||||
|
}
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = res.Body.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Read, so that the connection can be used again.
|
||||||
|
_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, maxAnswerBytes))
|
||||||
|
|
||||||
|
switch status := res.StatusCode; {
|
||||||
|
case status >= http.StatusOK && status < http.StatusMultipleChoices:
|
||||||
|
return nil
|
||||||
|
case status >= http.StatusBadRequest && status < http.StatusInternalServerError &&
|
||||||
|
status != http.StatusRequestTimeout && status != http.StatusTooManyRequests:
|
||||||
|
return fmt.Errorf("%w %s", errRefused, res.Status)
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("%w %s", errStatus, res.Status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// webhookMessage returns alert as JSON, for the webhook, and the headers
|
||||||
|
// sent with it: WebhookHeaders, and its Content-Type.
|
||||||
|
func (q *Queue) webhookMessage(alert *Alert) ([]byte, http.Header, error) {
|
||||||
|
body, err := json.Marshal(alert)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, fmt.Errorf("encode the alert: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
header := http.Header{}
|
||||||
|
maps.Copy(header, q.params.WebhookHeaders)
|
||||||
|
header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
return body, header, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// slackMessage returns alert as a message for a Slack incoming webhook,
|
||||||
|
// in JSON: its title in bold, then its text, with &, < and > escaped, as
|
||||||
|
// Slack asks, so that nothing in them is read as a link or a mention.
|
||||||
|
func slackMessage(alert *Alert) ([]byte, http.Header, error) {
|
||||||
|
escape := strings.NewReplacer("&", "&", "<", "<", ">", ">").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
@@ -0,0 +1,16 @@
|
|||||||
|
package alerts
|
||||||
|
|
||||||
|
import "net/http"
|
||||||
|
|
||||||
|
// QueueSize is the most alerts that wait to be sent to a destination.
|
||||||
|
const QueueSize = queueSize
|
||||||
|
|
||||||
|
// SetTransport has q's requests to the destination name go through
|
||||||
|
// transport instead of the network.
|
||||||
|
func (q *Queue) SetTransport(name string, transport http.RoundTripper) {
|
||||||
|
for _, d := range q.destinations {
|
||||||
|
if d.name == name {
|
||||||
|
d.httpClient.Transport = transport
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+14
-11
@@ -65,13 +65,16 @@ func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
|
|||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
|
|
||||||
limit := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
||||||
bans.Notes{Limit: 1000, Window: "minute"})
|
bans.Notes{Kind: "requests", Limit: 1000, Window: "minute"})
|
||||||
attack := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
|
byteLimit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.3/32"),
|
||||||
|
midnight(), bans.Notes{Kind: "bytes", Limit: 10 << 30, Window: "hour"})
|
||||||
|
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
|
||||||
bans.Notes{RuleID: "git-dir", Target: "path"})
|
bans.Notes{RuleID: "git-dir", Target: "path"})
|
||||||
|
|
||||||
for _, tc := range []struct{ got, want string }{
|
for _, tc := range []struct{ got, want string }{
|
||||||
{limit.Reason, "requests per minute over the limit of 1000"},
|
{limit.Reason, "requests per minute over the limit of 1000"},
|
||||||
|
{byteLimit.Reason, "bytes per hour over the limit of 10737418240"},
|
||||||
{attack.Reason, "matched the rule git-dir"},
|
{attack.Reason, "matched the rule git-dir"},
|
||||||
} {
|
} {
|
||||||
if tc.got != tc.want {
|
if tc.got != tc.want {
|
||||||
@@ -101,12 +104,12 @@ func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
|||||||
// kept, and counted among the earlier bans.
|
// kept, and counted among the earlier bans.
|
||||||
now := midnight().Add(30 * time.Minute)
|
now := midnight().Add(30 * time.Minute)
|
||||||
|
|
||||||
_, banned := ledger.Check(netblock.Addr(), now)
|
_, banned, _ := ledger.Check(netblock.Addr(), now)
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the lifted ban refuses")
|
t.Error("the lifted ban refuses")
|
||||||
}
|
}
|
||||||
|
|
||||||
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
if ban.Expires.Sub(ban.Start) != time.Hour ||
|
if ban.Expires.Sub(ban.Start) != time.Hour ||
|
||||||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
|
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
|
||||||
@@ -134,7 +137,7 @@ func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
|||||||
|
|
||||||
now := midnight().Add(2 * time.Hour)
|
now := midnight().Add(2 * time.Hour)
|
||||||
|
|
||||||
_, banned := ledger.Find(netblock.Addr(), now)
|
_, banned, _ := ledger.Find(netblock.Addr(), now)
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the lifted ban refuses")
|
t.Error("the lifted ban refuses")
|
||||||
}
|
}
|
||||||
@@ -145,7 +148,7 @@ func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// The next clear sign of attack bans for seven days, as a first does.
|
// The next clear sign of attack bans for seven days, as a first does.
|
||||||
ban := ledger.BanForAttack(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
|
||||||
if ban.Expires.Sub(ban.Start) != 7*day {
|
if ban.Expires.Sub(ban.Start) != 7*day {
|
||||||
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
|
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
|
||||||
}
|
}
|
||||||
@@ -155,7 +158,7 @@ func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
made := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
made, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
||||||
bans.Notes{})
|
bans.Notes{})
|
||||||
atStart := bans.Ban{
|
atStart := bans.Ban{
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
||||||
@@ -220,7 +223,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// It refuses once the ban for the limit has ended.
|
// It refuses once the ban for the limit has ended.
|
||||||
ban, banned := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
|
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
|
||||||
if !banned || ban != want {
|
if !banned || ban != want {
|
||||||
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
|
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
|
||||||
ban, banned, want)
|
ban, banned, want)
|
||||||
@@ -261,11 +264,11 @@ func TestLiftLiftsEveryActiveBanCoveringTheClient(t *testing.T) {
|
|||||||
|
|
||||||
wantChanged(t, ledger, true)
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
if _, banned := ledger.Check(client, now); banned {
|
if _, banned, _ := ledger.Check(client, now); banned {
|
||||||
t.Error("the client is still banned")
|
t.Error("the client is still banned")
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, banned := ledger.Check(other.Addr(), now); !banned {
|
if _, banned, _ := ledger.Check(other.Addr(), now); !banned {
|
||||||
t.Error("the other client's ban was lifted")
|
t.Error("the other client's ban was lifted")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+139
-42
@@ -1,8 +1,8 @@
|
|||||||
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
||||||
// netblocks of clients that break a rate limit or show a clear sign of
|
// netblocks of clients that break a rate limit or a byte limit or show a
|
||||||
// attack, and those an admin makes, with their notes, as the "Bans"
|
// clear sign of attack, and those an admin makes, with their notes, as
|
||||||
// section of SPEC.md describes. The bans are kept in memory, and written
|
// the "Bans" section of SPEC.md describes. The bans are kept in memory,
|
||||||
// to bans.json and read from it by the state package.
|
// and written to bans.json and read from it by the state package.
|
||||||
package bans
|
package bans
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -91,22 +91,35 @@ func (b Ban) ActiveAt(now time.Time) bool {
|
|||||||
//
|
//
|
||||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
type Notes struct {
|
type Notes struct {
|
||||||
// Country is the client's country, when it was looked up.
|
// ASN, ASName and Country are the client's AS number, AS name and
|
||||||
|
// country, when they were looked up: when the request that caused the
|
||||||
|
// ban was made, or when GeoJS answered about the client afterwards.
|
||||||
|
ASN string `json:"asn"`
|
||||||
|
ASName string `json:"as_name"`
|
||||||
Country string `json:"country"`
|
Country string `json:"country"`
|
||||||
// Limit, Window and Count are, for a ban for a broken limit, the limit
|
// Kind, Limit, Window and Count are, for a ban for a broken limit,
|
||||||
// that was broken, its window, "minute", "hour" or "day", and the
|
// what the limit was on, "requests" for a rate limit or "bytes" for a
|
||||||
// count reached: the client's requests in the window, the one that
|
// byte limit, the limit that was broken, its window, "minute", "hour"
|
||||||
// broke the limit included. These are the requests that counted
|
// or "day", and the count reached: the client's requests, or bytes, in
|
||||||
// toward the ban, and the window is the time over which they came.
|
// the window, those of the request that broke the limit included.
|
||||||
|
// These are what counted toward the ban, and the window is the time
|
||||||
|
// over which they came.
|
||||||
|
Kind string `json:"kind,omitempty"`
|
||||||
Limit int64 `json:"limit,omitempty"`
|
Limit int64 `json:"limit,omitempty"`
|
||||||
Window string `json:"window,omitempty"`
|
Window string `json:"window,omitempty"`
|
||||||
Count float64 `json:"count,omitempty"`
|
Count float64 `json:"count,omitempty"`
|
||||||
|
// LimitPercent and LimitPercentSetting are, for a ban for a limit a
|
||||||
|
// biased threshold lowered, the client's percentage of that kind of
|
||||||
|
// limit, of which Limit is the result, and the setting that gave it.
|
||||||
|
// Both are left out for a limit that was not lowered.
|
||||||
|
LimitPercent *int64 `json:"limit_percent,omitempty"`
|
||||||
|
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
|
||||||
// RuleID and Target are, for a ban for a clear sign of attack, the id
|
// RuleID and Target are, for a ban for a clear sign of attack, the id
|
||||||
// of the rule file rule that matched, and its target.
|
// of the rule file rule that matched, and its target.
|
||||||
RuleID string `json:"rule_id,omitempty"`
|
RuleID string `json:"rule_id,omitempty"`
|
||||||
Target string `json:"target,omitempty"`
|
Target string `json:"target,omitempty"`
|
||||||
// Request is the request that broke the limit, or that was the clear
|
// Request is the request that broke the limit, or whose bytes broke
|
||||||
// sign of attack.
|
// it, or that was the clear sign of attack.
|
||||||
Request Request `json:"request"`
|
Request Request `json:"request"`
|
||||||
// Requests is how many requests the netblock has sent since it was
|
// Requests is how many requests the netblock has sent since it was
|
||||||
// first seen, and Refused how many of them the ban has refused so
|
// first seen, and Refused how many of them the ban has refused so
|
||||||
@@ -194,39 +207,43 @@ func (l *Ledger) Changed() <-chan struct{} {
|
|||||||
// a ban on a netblock client is in is active, and returns that ban, with
|
// a ban on a netblock client is in is active, and returns that ban, with
|
||||||
// the request counted among those it refused. A ban for a clear sign of
|
// the request counted among those it refused. A ban for a clear sign of
|
||||||
// attack is made permanent by the request: the netblock is malicious.
|
// attack is made permanent by the request: the netblock is malicious.
|
||||||
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
|
// The last result reports whether the request made the ban permanent.
|
||||||
|
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool, bool) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
ban := l.active(client, now)
|
ban := l.active(client, now)
|
||||||
if ban == nil {
|
if ban == nil {
|
||||||
return Ban{}, false
|
return Ban{}, false, false
|
||||||
}
|
}
|
||||||
|
|
||||||
ban.Notes.Requests++
|
ban.Notes.Requests++
|
||||||
ban.Notes.Refused++
|
ban.Notes.Refused++
|
||||||
|
|
||||||
if ban.Cause == CauseAttack && !ban.Permanent() {
|
madePermanent := ban.Cause == CauseAttack && !ban.Permanent()
|
||||||
|
if madePermanent {
|
||||||
ban.Expires = time.Time{}
|
ban.Expires = time.Time{}
|
||||||
|
|
||||||
l.markChanged()
|
l.markChanged()
|
||||||
}
|
}
|
||||||
|
|
||||||
return *ban, true
|
return *ban, true, madePermanent
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find is Check without counting the request among those the ban
|
// Find is Check without counting the request among those the ban
|
||||||
// refused: in observe mode a ban refuses nothing.
|
// refused, and without making the ban permanent: in observe mode a ban
|
||||||
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
|
// refuses nothing. The last result reports whether Check would have made
|
||||||
|
// the ban permanent.
|
||||||
|
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
ban := l.active(client, now)
|
ban := l.active(client, now)
|
||||||
if ban == nil {
|
if ban == nil {
|
||||||
return Ban{}, false
|
return Ban{}, false, false
|
||||||
}
|
}
|
||||||
|
|
||||||
return *ban, true
|
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent()
|
||||||
}
|
}
|
||||||
|
|
||||||
// activeBan returns the ban in bans, a netblock's bans oldest first, that
|
// activeBan returns the ban in bans, a netblock's bans oldest first, that
|
||||||
@@ -244,28 +261,80 @@ func activeBan(bans []Ban, now time.Time) *Ban {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// BanForLimit bans netblock at now for a broken limit, with notes, and
|
// BanForLimit bans netblock at now for a broken limit, with notes, and
|
||||||
// returns the ban. A first ban lasts LimitBanDuration. A ban made within
|
// returns the ban, and true. A first ban lasts LimitBanDuration. A ban
|
||||||
// LimitBanRepeatWindow after the netblock's ban that ended last, other
|
// made within LimitBanRepeatWindow after the netblock's ban that ended
|
||||||
// than one for a clear sign of attack or a lifted one, lasts repeatFactor
|
// last, other than one for a clear sign of attack or a lifted one, lasts
|
||||||
// times as long as that one. A ban that would be longer than
|
// repeatFactor times as long as that one. A ban that would be longer
|
||||||
// MaxBanDuration is permanent instead. If a ban on netblock is still
|
// than MaxBanDuration is permanent instead. If a ban on netblock is still
|
||||||
// active, as when two of its requests break a limit at once, that ban is
|
// active, as when two of its requests break a limit at once, that ban is
|
||||||
// returned and no other is made. The ledger fills in the notes' Refused
|
// returned with false, and no other is made. The ledger fills in the
|
||||||
// and EarlierBans itself, and gives the ban the reason "requests per
|
// notes' Refused and EarlierBans itself, and gives the ban the reason
|
||||||
// <Window> over the limit of <Limit>", from the notes.
|
// "<Kind> per <Window> over the limit of <Limit>", from the notes, such
|
||||||
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
|
// as "requests per minute over the limit of 1000".
|
||||||
reason := fmt.Sprintf("requests per %s over the limit of %d",
|
func (l *Ledger) BanForLimit(
|
||||||
notes.Window, notes.Limit)
|
netblock netip.Prefix, now time.Time, notes Notes,
|
||||||
|
) (Ban, bool) {
|
||||||
|
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
|
||||||
|
}
|
||||||
|
|
||||||
return l.ban(netblock, now, CauseLimit, reason, notes)
|
// WouldBanForLimit returns what BanForLimit would, without making the ban:
|
||||||
|
// what observe mode would have done.
|
||||||
|
func (l *Ledger) WouldBanForLimit(
|
||||||
|
netblock netip.Prefix, now time.Time, notes Notes,
|
||||||
|
) (Ban, bool) {
|
||||||
|
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
|
||||||
}
|
}
|
||||||
|
|
||||||
// BanForAttack bans netblock at now for a clear sign of attack, with
|
// BanForAttack bans netblock at now for a clear sign of attack, with
|
||||||
// notes, and returns the ban, as BanForLimit does. A first ban lasts
|
// notes, and returns the ban, and whether it made it, as BanForLimit
|
||||||
// AttackBanDuration; once the netblock has had one that was not lifted,
|
// does. A first ban lasts AttackBanDuration; once the netblock has had
|
||||||
// the next is permanent. Its reason is "matched the rule <RuleID>".
|
// one that was not lifted, the next is permanent. Its reason is "matched
|
||||||
func (l *Ledger) BanForAttack(netblock netip.Prefix, now time.Time, notes Notes) Ban {
|
// the rule <RuleID>".
|
||||||
return l.ban(netblock, now, CauseAttack, "matched the rule "+notes.RuleID, notes)
|
func (l *Ledger) BanForAttack(
|
||||||
|
netblock netip.Prefix, now time.Time, notes Notes,
|
||||||
|
) (Ban, bool) {
|
||||||
|
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WouldBanForAttack returns what BanForAttack would, without making the
|
||||||
|
// ban: what observe mode would have done.
|
||||||
|
func (l *Ledger) WouldBanForAttack(
|
||||||
|
netblock netip.Prefix, now time.Time, notes Notes,
|
||||||
|
) (Ban, bool) {
|
||||||
|
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit
|
||||||
|
// or CauseAttack, made at now would be permanent, as BanForLimit or
|
||||||
|
// BanForAttack would make it. It works out nothing else of the ban.
|
||||||
|
func (l *Ledger) WouldBePermanent(
|
||||||
|
netblock netip.Prefix, now time.Time, cause string,
|
||||||
|
) bool {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
var held []Ban
|
||||||
|
if bans, found := l.netblocks.Peek(netblock); found {
|
||||||
|
held = *bans
|
||||||
|
}
|
||||||
|
|
||||||
|
if cause == CauseAttack {
|
||||||
|
return l.attackExpiry(held, now).IsZero()
|
||||||
|
}
|
||||||
|
|
||||||
|
return l.limitExpiry(held, now).IsZero()
|
||||||
|
}
|
||||||
|
|
||||||
|
// limitReason is the reason of a ban for a broken limit, with notes.
|
||||||
|
func limitReason(notes Notes) string {
|
||||||
|
return fmt.Sprintf("%s per %s over the limit of %d",
|
||||||
|
notes.Kind, notes.Window, notes.Limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
// attackReason is the reason of a ban for a clear sign of attack, with
|
||||||
|
// notes.
|
||||||
|
func attackReason(notes Notes) string {
|
||||||
|
return "matched the rule " + notes.RuleID
|
||||||
}
|
}
|
||||||
|
|
||||||
// BanForAdmin bans netblock at now for an admin, with reason, until
|
// BanForAdmin bans netblock at now for an admin, with reason, until
|
||||||
@@ -357,6 +426,28 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
|
|||||||
return slices.Clone(*bans)
|
return slices.Clone(*bans)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddLookup gives the notes of netblock's bans that have no AS number, AS
|
||||||
|
// name or country yet those of a client in it, as the lookup answered
|
||||||
|
// about it. It is not a request from netblock, and leaves when it was last
|
||||||
|
// seen unchanged. It does not have bans.json written at once: the notes
|
||||||
|
// are written with its next write, as the counts in them are.
|
||||||
|
func (l *Ledger) AddLookup(netblock netip.Prefix, asn, asName, country string) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
bans, found := l.netblocks.Peek(netblock)
|
||||||
|
if !found {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range *bans {
|
||||||
|
notes := &(*bans)[i].Notes
|
||||||
|
if notes.ASN == "" && notes.ASName == "" && notes.Country == "" {
|
||||||
|
notes.ASN, notes.ASName, notes.Country = asn, asName, country
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Made returns how many bans for cause have been made since the start:
|
// Made returns how many bans for cause have been made since the start:
|
||||||
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
|
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
|
||||||
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
|
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
|
||||||
@@ -483,10 +574,12 @@ func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ban bans netblock at now for cause, with reason and notes, as
|
// ban bans netblock at now for cause, with reason and notes, as
|
||||||
// BanForLimit and BanForAttack describe, and returns the ban.
|
// BanForLimit and BanForAttack describe, and returns the ban, and whether
|
||||||
|
// it made it. Unless keep is true, the ban is not made, only returned: it
|
||||||
|
// is the ban that would have been made.
|
||||||
func (l *Ledger) ban(
|
func (l *Ledger) ban(
|
||||||
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes,
|
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool,
|
||||||
) Ban {
|
) (Ban, bool) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
@@ -497,7 +590,7 @@ func (l *Ledger) ban(
|
|||||||
if found {
|
if found {
|
||||||
active := activeBan(*bans, now)
|
active := activeBan(*bans, now)
|
||||||
if active != nil {
|
if active != nil {
|
||||||
return *active
|
return *active, false
|
||||||
}
|
}
|
||||||
|
|
||||||
held = *bans
|
held = *bans
|
||||||
@@ -513,11 +606,15 @@ func (l *Ledger) ban(
|
|||||||
ban.Expires = l.limitExpiry(held, now)
|
ban.Expires = l.limitExpiry(held, now)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !keep {
|
||||||
|
return ban, true
|
||||||
|
}
|
||||||
|
|
||||||
l.add(ban)
|
l.add(ban)
|
||||||
l.made[cause]++
|
l.made[cause]++
|
||||||
l.markChanged()
|
l.markChanged()
|
||||||
|
|
||||||
return ban
|
return ban, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// earlierBans returns how many bans a netblock with the bans held, oldest
|
// earlierBans returns how many bans a netblock with the bans held, oldest
|
||||||
|
|||||||
+179
-44
@@ -21,7 +21,7 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
|
|||||||
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
|
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
|
||||||
// 81 hours.
|
// 81 hours.
|
||||||
for i, hours := range []int{1, 3, 9, 27, 81} {
|
for i, hours := range []int{1, 3, 9, 27, 81} {
|
||||||
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
|
|
||||||
length := time.Duration(hours) * time.Hour
|
length := time.Duration(hours) * time.Hour
|
||||||
if !ban.Expires.Equal(now.Add(length)) ||
|
if !ban.Expires.Equal(now.Add(length)) ||
|
||||||
@@ -35,12 +35,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
|
|||||||
|
|
||||||
// The sixth would last 243 hours, more than seven days: it is
|
// The sixth would last 243 hours, more than seven days: it is
|
||||||
// permanent, and never ends.
|
// permanent, and never ends.
|
||||||
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
if !ban.Permanent() {
|
if !ban.Permanent() {
|
||||||
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
|
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day))
|
_, banned, _ := ledger.Check(netblock.Addr(), now.Add(100*365*day))
|
||||||
if !banned {
|
if !banned {
|
||||||
t.Error("a permanent ban ended")
|
t.Error("a permanent ban ended")
|
||||||
}
|
}
|
||||||
@@ -64,8 +64,8 @@ func TestRepeatWindowRunsOut(t *testing.T) {
|
|||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
|
second, _ := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
|
||||||
|
|
||||||
if second.Expires.Sub(second.Start) != tc.want ||
|
if second.Expires.Sub(second.Start) != tc.want ||
|
||||||
second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
||||||
@@ -83,7 +83,7 @@ func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
|
|||||||
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
|
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
|
||||||
ledger := bans.New(rules)
|
ledger := bans.New(rules)
|
||||||
|
|
||||||
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
ban, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
||||||
bans.Notes{})
|
bans.Notes{})
|
||||||
if !ban.Permanent() {
|
if !ban.Permanent() {
|
||||||
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
|
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
|
||||||
@@ -103,7 +103,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
|
|||||||
now := midnight()
|
now := midnight()
|
||||||
|
|
||||||
for i := range 14 {
|
for i := range 14 {
|
||||||
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
if !ban.Expires.After(ban.Start) {
|
if !ban.Expires.After(ban.Start) {
|
||||||
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
|
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
|
||||||
}
|
}
|
||||||
@@ -111,7 +111,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
|
|||||||
now = ban.Expires
|
now = ban.Expires
|
||||||
}
|
}
|
||||||
|
|
||||||
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
if !ban.Permanent() {
|
if !ban.Permanent() {
|
||||||
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
|
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
|
||||||
}
|
}
|
||||||
@@ -123,12 +123,22 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
|
|||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
first, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
|
if !made {
|
||||||
|
t.Error("the first ban was not made")
|
||||||
|
}
|
||||||
|
|
||||||
if again != first || len(ledger.Bans(netblock)) != 1 {
|
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
|
||||||
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
|
|
||||||
again, len(ledger.Bans(netblock)), first)
|
if made || again != first || len(ledger.Bans(netblock)) != 1 {
|
||||||
|
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
|
||||||
|
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
|
||||||
|
}
|
||||||
|
|
||||||
|
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
|
||||||
|
if made || again != first {
|
||||||
|
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
|
||||||
|
again, made, first)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -137,21 +147,21 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
|
|||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
||||||
|
|
||||||
for range 3 {
|
for range 3 {
|
||||||
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
got, banned, _ := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
||||||
if !banned || got.Start != ban.Start {
|
if !banned || got.Start != ban.Start {
|
||||||
t.Fatalf("check during the ban gives %+v and %t", got, banned)
|
t.Fatalf("check during the ban gives %+v and %t", got, banned)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
|
_, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("another netblock is banned")
|
t.Error("another netblock is banned")
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned = ledger.Check(netblock.Addr(), ban.Expires)
|
_, banned, _ = ledger.Check(netblock.Addr(), ban.Expires)
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the ban did not end")
|
t.Error("the ban did not end")
|
||||||
}
|
}
|
||||||
@@ -169,14 +179,14 @@ func TestFindCountsNothing(t *testing.T) {
|
|||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
||||||
|
|
||||||
got, banned := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
got, banned, _ := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
||||||
if !banned || got != ban {
|
if !banned || got != ban {
|
||||||
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
|
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned = ledger.Find(netblock.Addr(), ban.Expires)
|
_, banned, _ = ledger.Find(netblock.Addr(), ban.Expires)
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the ban did not end")
|
t.Error("the ban did not end")
|
||||||
}
|
}
|
||||||
@@ -198,7 +208,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
|
|||||||
d := netip.MustParsePrefix("2001:db8::/64")
|
d := netip.MustParsePrefix("2001:db8::/64")
|
||||||
now := midnight()
|
now := midnight()
|
||||||
|
|
||||||
first := ledger.BanForLimit(a, now, bans.Notes{})
|
first, _ := ledger.BanForLimit(a, now, bans.Notes{})
|
||||||
ledger.BanForLimit(b, now, bans.Notes{})
|
ledger.BanForLimit(b, now, bans.Notes{})
|
||||||
ledger.BanForLimit(c, now, bans.Notes{})
|
ledger.BanForLimit(c, now, bans.Notes{})
|
||||||
|
|
||||||
@@ -233,8 +243,8 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
|
|||||||
ledger := bans.New(rules)
|
ledger := bans.New(rules)
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
second := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
|
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
|
||||||
|
|
||||||
held := ledger.Bans(netblock)
|
held := ledger.Bans(netblock)
|
||||||
if len(held) != 1 || held[0] != second ||
|
if len(held) != 1 || held[0] != second ||
|
||||||
@@ -251,7 +261,7 @@ func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
|
|||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
notes := bans.Notes{RuleID: "env-file", Target: "path"}
|
notes := bans.Notes{RuleID: "env-file", Target: "path"}
|
||||||
|
|
||||||
ban := ledger.BanForAttack(netblock, midnight(), notes)
|
ban, _ := ledger.BanForAttack(netblock, midnight(), notes)
|
||||||
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
|
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
|
||||||
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
|
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
|
||||||
ledger.Made(bans.CauseLimit) != 0 {
|
ledger.Made(bans.CauseLimit) != 0 {
|
||||||
@@ -262,23 +272,31 @@ func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
|
|||||||
|
|
||||||
wantChanged(t, ledger, true)
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
// In observe mode the ban refuses nothing, and stays as it is.
|
// In observe mode the ban refuses nothing, and stays as it is, while
|
||||||
got, _ := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
|
// Find tells that the request would have made it permanent.
|
||||||
if got.Permanent() {
|
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
|
||||||
t.Fatal("a request found under the ban made it permanent")
|
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent {
|
||||||
|
t.Fatalf("a request found under the ban left it %+v, would have made it "+
|
||||||
|
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
|
||||||
}
|
}
|
||||||
|
|
||||||
// A request it refuses makes it permanent, and bans.json due.
|
wantChanged(t, ledger, false)
|
||||||
got, _ = ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
|
|
||||||
if !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
|
// A request it refuses makes it permanent, says so, and makes
|
||||||
t.Fatalf("after a request during the ban, it is %+v, want it permanent", got)
|
// bans.json due.
|
||||||
|
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
|
||||||
|
if !madePermanent || !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
|
||||||
|
t.Fatalf("after a request during the ban, it is %+v, made permanent %t, "+
|
||||||
|
"want it made permanent", got, madePermanent)
|
||||||
}
|
}
|
||||||
|
|
||||||
wantChanged(t, ledger, true)
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
_, banned := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
|
// The next request finds it permanent already.
|
||||||
if !banned {
|
_, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
|
||||||
t.Error("the permanent ban ended")
|
if !banned || madePermanent {
|
||||||
|
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
|
||||||
|
"want banned by the permanent ban", banned, madePermanent)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -289,8 +307,8 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
|
|||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
// A ban for a broken limit before does not count.
|
// A ban for a broken limit before does not count.
|
||||||
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
second := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
|
second, _ := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
|
||||||
|
|
||||||
if second.Expires.Sub(second.Start) != 7*day {
|
if second.Expires.Sub(second.Start) != 7*day {
|
||||||
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
|
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
|
||||||
@@ -299,14 +317,14 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
|
|||||||
|
|
||||||
// Once that has run out without a request, the netblock is served, and
|
// Once that has run out without a request, the netblock is served, and
|
||||||
// its next clear sign of attack bans it for good.
|
// its next clear sign of attack bans it for good.
|
||||||
_, banned := ledger.Check(netblock.Addr(), second.Expires)
|
_, banned, _ := ledger.Check(netblock.Addr(), second.Expires)
|
||||||
if banned {
|
if banned {
|
||||||
t.Fatal("the ban did not end")
|
t.Fatal("the ban did not end")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Its notes show the earlier ban for an attack that makes it permanent,
|
// Its notes show the earlier ban for an attack that makes it permanent,
|
||||||
// beside the one for a limit.
|
// beside the one for a limit.
|
||||||
third := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
|
third, _ := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
|
||||||
if !third.Permanent() ||
|
if !third.Permanent() ||
|
||||||
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
|
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
|
||||||
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
|
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
|
||||||
@@ -314,6 +332,84 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
|
// While the first ban lasts, none would be made.
|
||||||
|
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
|
||||||
|
if would || during != first {
|
||||||
|
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
|
||||||
|
would, during, first)
|
||||||
|
}
|
||||||
|
|
||||||
|
// As it ends, a clear sign of attack would ban for seven days, and a
|
||||||
|
// limit broken again for three hours, but neither is made.
|
||||||
|
limitNotes := bans.Notes{Kind: "requests", Limit: 1, Window: "minute"}
|
||||||
|
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
|
||||||
|
bans.Notes{RuleID: "git-dir"})
|
||||||
|
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
|
||||||
|
|
||||||
|
if !wouldAttack || !attack.Expires.Equal(first.Expires.Add(7*day)) ||
|
||||||
|
attack.Reason != "matched the rule git-dir" || !wouldLimit ||
|
||||||
|
!limit.Expires.Equal(first.Expires.Add(3*time.Hour)) ||
|
||||||
|
limit.Reason != "requests per minute over the limit of 1" {
|
||||||
|
t.Errorf("would ban with %+v and %+v, want seven days for the attack and "+
|
||||||
|
"three hours for the limit", attack, limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(ledger.Bans(netblock)) != 1 || ledger.Made(bans.CauseLimit) != 1 ||
|
||||||
|
ledger.Made(bans.CauseAttack) != 0 {
|
||||||
|
t.Errorf("the ledger holds %+v, want the first ban alone", ledger.Bans(netblock))
|
||||||
|
}
|
||||||
|
|
||||||
|
wantChanged(t, ledger, false)
|
||||||
|
|
||||||
|
// The ban made is the one that would have been.
|
||||||
|
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
|
||||||
|
if made != limit {
|
||||||
|
t.Errorf("the ban made is %+v, want %+v", made, limit)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWouldBePermanentAnswersAsTheBanWouldBeMade(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
now := midnight()
|
||||||
|
|
||||||
|
// Five bans for a limit in a row, of 1, 3, 9, 27 and 81 hours, are not
|
||||||
|
// permanent. The sixth, of 243 hours, would be, while a first ban for
|
||||||
|
// an attack would not.
|
||||||
|
for i := range 5 {
|
||||||
|
if ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
|
||||||
|
t.Fatalf("ban %d for a limit would be permanent", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
|
now = ban.Expires
|
||||||
|
}
|
||||||
|
|
||||||
|
if !ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
|
||||||
|
t.Error("the sixth ban for a limit would not be permanent")
|
||||||
|
}
|
||||||
|
|
||||||
|
if ledger.WouldBePermanent(netblock, now, bans.CauseAttack) {
|
||||||
|
t.Error("a first ban for an attack would be permanent")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Once a first ban for an attack has ended, the next would be permanent.
|
||||||
|
attack, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
|
||||||
|
if !ledger.WouldBePermanent(netblock, attack.Expires, bans.CauseAttack) {
|
||||||
|
t.Error("a second ban for an attack would not be permanent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
|
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -322,16 +418,16 @@ func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
|
|||||||
|
|
||||||
// Three times the seven days would be permanent; a limit broken as the
|
// Three times the seven days would be permanent; a limit broken as the
|
||||||
// ban for an attack ends bans for an hour, as a first broken limit does.
|
// ban for an attack ends bans for an hour, as a first broken limit does.
|
||||||
attack := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
|
attack, _ := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
|
||||||
limit := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
|
limit, _ := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
|
||||||
|
|
||||||
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
|
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
|
||||||
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
|
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
|
||||||
}
|
}
|
||||||
|
|
||||||
// And a request during the ban for a limit leaves it as it is.
|
// And a request during the ban for a limit leaves it as it is.
|
||||||
got, _ := ledger.Check(netblock.Addr(), limit.Start)
|
got, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start)
|
||||||
if got.Permanent() {
|
if got.Permanent() || madePermanent {
|
||||||
t.Error("a request during a ban for a limit made it permanent")
|
t.Error("a request during a ban for a limit made it permanent")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -346,7 +442,7 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
|
|||||||
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
|
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
|
||||||
}
|
}
|
||||||
|
|
||||||
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
|
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
|
||||||
|
|
||||||
cut := long[:256]
|
cut := long[:256]
|
||||||
want := bans.Request{
|
want := bans.Request{
|
||||||
@@ -358,6 +454,45 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLookupFillsTheNotesOfTheNetblocksBansWithoutOne(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
other := netip.MustParsePrefix("198.51.100.7/32")
|
||||||
|
|
||||||
|
// A ban made with the client's lookup, one made before it came, after
|
||||||
|
// the first ended, and one on another netblock.
|
||||||
|
ledger.BanForLimit(netblock, midnight(), bans.Notes{
|
||||||
|
ASN: "AS64497", ASName: "Other Net", Country: "FR",
|
||||||
|
})
|
||||||
|
ledger.BanForLimit(netblock, midnight().Add(time.Hour), bans.Notes{})
|
||||||
|
ledger.BanForLimit(other, midnight(), bans.Notes{})
|
||||||
|
|
||||||
|
ledger.AddLookup(netblock, "AS64496", "Example Net", "DE")
|
||||||
|
|
||||||
|
held := ledger.Bans(netblock)
|
||||||
|
if len(held) != 2 {
|
||||||
|
t.Fatalf("%s has %d bans, want 2", netblock, len(held))
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, want := range []bans.Notes{
|
||||||
|
{ASN: "AS64497", ASName: "Other Net", Country: "FR"},
|
||||||
|
{ASN: "AS64496", ASName: "Example Net", Country: "DE"},
|
||||||
|
} {
|
||||||
|
got := held[i].Notes
|
||||||
|
if got.ASN != want.ASN || got.ASName != want.ASName || got.Country != want.Country {
|
||||||
|
t.Errorf("ban %d's notes give %q, %q and %q, want %q, %q and %q", i+1,
|
||||||
|
got.ASN, got.ASName, got.Country, want.ASN, want.ASName, want.Country)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if notes := ledger.Bans(other)[0].Notes; notes.ASN != "" || notes.Country != "" {
|
||||||
|
t.Errorf("the ban on %s has %q and %q, want neither",
|
||||||
|
other, notes.ASN, notes.Country)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// defaultRules are the rules at the settings' defaults.
|
// defaultRules are the rules at the settings' defaults.
|
||||||
func defaultRules() bans.Rules {
|
func defaultRules() bans.Rules {
|
||||||
return bans.Rules{
|
return bans.Rules{
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
|
|||||||
high := netip.MustParsePrefix("203.0.113.10/32")
|
high := netip.MustParsePrefix("203.0.113.10/32")
|
||||||
low := netip.MustParsePrefix("203.0.113.9/32")
|
low := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
first := ledger.BanForLimit(v6, midnight(), bans.Notes{})
|
first, _ := ledger.BanForLimit(v6, midnight(), bans.Notes{})
|
||||||
ledger.BanForLimit(high, midnight(), bans.Notes{})
|
ledger.BanForLimit(high, midnight(), bans.Notes{})
|
||||||
ledger.BanForLimit(low, midnight(), bans.Notes{})
|
ledger.BanForLimit(low, midnight(), bans.Notes{})
|
||||||
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
|
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
|
||||||
@@ -68,7 +68,7 @@ func TestLoadedBansCarryOn(t *testing.T) {
|
|||||||
|
|
||||||
before := bans.New(defaultRules())
|
before := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
|
ban, _ := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
|
||||||
|
|
||||||
// Loaded into a new ledger, as across a restart, the ban still refuses
|
// Loaded into a new ledger, as across a restart, the ban still refuses
|
||||||
// while it lasts, and once it has ended a broken limit bans for three
|
// while it lasts, and once it has ended a broken limit bans for three
|
||||||
@@ -76,12 +76,12 @@ func TestLoadedBansCarryOn(t *testing.T) {
|
|||||||
after := bans.New(defaultRules())
|
after := bans.New(defaultRules())
|
||||||
after.Load(before.Snapshot())
|
after.Load(before.Snapshot())
|
||||||
|
|
||||||
_, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
|
_, banned, _ := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
|
||||||
if !banned {
|
if !banned {
|
||||||
t.Error("the loaded ban does not refuse")
|
t.Error("the loaded ban does not refuse")
|
||||||
}
|
}
|
||||||
|
|
||||||
again := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
|
again, _ := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
|
||||||
if again.Expires.Sub(again.Start) != 3*time.Hour ||
|
if again.Expires.Sub(again.Start) != 3*time.Hour ||
|
||||||
again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
|
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
|
||||||
@@ -111,7 +111,7 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
|
|||||||
"198.51.100.7": true,
|
"198.51.100.7": true,
|
||||||
"198.51.100.8": false,
|
"198.51.100.8": false,
|
||||||
} {
|
} {
|
||||||
_, banned := ledger.Check(netip.MustParseAddr(client), midnight())
|
_, banned, _ := ledger.Check(netip.MustParseAddr(client), midnight())
|
||||||
if banned != want {
|
if banned != want {
|
||||||
t.Errorf("%s is refused: %t, want %t", client, banned, want)
|
t.Errorf("%s is refused: %t, want %t", client, banned, want)
|
||||||
}
|
}
|
||||||
@@ -150,19 +150,19 @@ func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
|
|||||||
now := midnight().Add(2 * time.Hour)
|
now := midnight().Add(2 * time.Hour)
|
||||||
client := netip.MustParseAddr("203.0.113.9")
|
client := netip.MustParseAddr("203.0.113.9")
|
||||||
|
|
||||||
ban, banned := ledger.Find(client, now)
|
ban, banned, _ := ledger.Find(client, now)
|
||||||
if !banned || !ban.Permanent() {
|
if !banned || !ban.Permanent() {
|
||||||
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
|
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
|
||||||
}
|
}
|
||||||
|
|
||||||
ban, banned = ledger.Check(client, now)
|
ban, banned, _ = ledger.Check(client, now)
|
||||||
if !banned || !ban.Permanent() {
|
if !banned || !ban.Permanent() {
|
||||||
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
|
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
|
||||||
banned, ban)
|
banned, ban)
|
||||||
}
|
}
|
||||||
|
|
||||||
// A limit broken now makes no shorter ban over the permanent one.
|
// A limit broken now makes no shorter ban over the permanent one.
|
||||||
ban = ledger.BanForLimit(netblock, now, bans.Notes{})
|
ban, _ = ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
|
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
|
||||||
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
|
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
|
||||||
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
|
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
|
||||||
@@ -194,7 +194,7 @@ func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
|
|||||||
// Once both have ended, a limit broken within the repeat window bans
|
// Once both have ended, a limit broken within the repeat window bans
|
||||||
// for three times the 9 hours, and the notes count the two bans
|
// for three times the 9 hours, and the notes count the two bans
|
||||||
// before the 9-hour one and it, for a limit, and the admin's.
|
// before the 9-hour one and it, for a limit, and the admin's.
|
||||||
ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
|
||||||
if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
|
if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
|
||||||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
|
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
|
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
|
||||||
@@ -255,15 +255,15 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
|
|||||||
// bans.json is taken in, that ban is lifted.
|
// bans.json is taken in, that ban is lifted.
|
||||||
ledger.Load([]bans.Ban{kept})
|
ledger.Load([]bans.Ban{kept})
|
||||||
|
|
||||||
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
|
_, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("a ban left out of the second load still refuses")
|
t.Error("a ban left out of the second load still refuses")
|
||||||
}
|
}
|
||||||
|
|
||||||
// The ledger holds one ban, so it makes two more without dropping any.
|
// The ledger holds one ban, so it makes two more without dropping any.
|
||||||
first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
|
first, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
|
||||||
bans.Notes{})
|
bans.Notes{})
|
||||||
second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
|
second, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
|
||||||
bans.Notes{})
|
bans.Notes{})
|
||||||
|
|
||||||
want := []bans.Ban{first, second, kept}
|
want := []bans.Ban{first, second, kept}
|
||||||
|
|||||||
+494
-28
@@ -20,8 +20,10 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -32,9 +34,10 @@ type Config struct {
|
|||||||
ListenAddr string
|
ListenAddr string
|
||||||
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
|
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
|
||||||
UpstreamURL *url.URL
|
UpstreamURL *url.URL
|
||||||
// InstanceName is the name each request log line gives as instance
|
// InstanceName is the name every log line and alert gives as instance,
|
||||||
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets
|
// and every metric carries as its label instance (SWWAF_INSTANCE_NAME),
|
||||||
// to the first 12 characters of the container's id.
|
// by default the host's name, which docker sets to the first 12
|
||||||
|
// characters of the container's id.
|
||||||
InstanceName string
|
InstanceName string
|
||||||
// Observe is true in observe mode, when SWWAF_MODE is observe rather
|
// Observe is true in observe mode, when SWWAF_MODE is observe rather
|
||||||
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country
|
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country
|
||||||
@@ -88,6 +91,28 @@ type Config struct {
|
|||||||
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
|
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
|
||||||
// Each starts with /.
|
// Each starts with /.
|
||||||
RateLimitExemptPaths []string
|
RateLimitExemptPaths []string
|
||||||
|
// BytesLimitPerMinute, BytesLimitPerHour and BytesLimitPerDay are the
|
||||||
|
// most bytes a client's requests may carry in a minute, an hour and a
|
||||||
|
// day (SWWAF_BYTES_LIMIT_PER_MINUTE, SWWAF_BYTES_LIMIT_PER_HOUR and
|
||||||
|
// SWWAF_BYTES_LIMIT_PER_DAY). BytesCount is which body bytes count
|
||||||
|
// toward them (SWWAF_BYTES_COUNT): response, request or both.
|
||||||
|
BytesLimitPerMinute int64
|
||||||
|
BytesLimitPerHour int64
|
||||||
|
BytesLimitPerDay int64
|
||||||
|
BytesCount string
|
||||||
|
// LookupSource is where each client's AS number and country are
|
||||||
|
// looked up (SWWAF_LOOKUP_SOURCE): geojs, file, or off for nowhere.
|
||||||
|
// LookupDBPath is the lookup database, the IPinfo Lite file looked up
|
||||||
|
// in while LookupSource is file (SWWAF_LOOKUP_DB_PATH), and "" for any
|
||||||
|
// other source. A request waits up to LookupTimeout for its client's
|
||||||
|
// first answer from GeoJS while a setting needs it
|
||||||
|
// (SWWAF_LOOKUP_TIMEOUT), which cannot be off. AddLookupHeaders is true
|
||||||
|
// when the app is passed the client's AS number and country in headers
|
||||||
|
// (SWWAF_ADD_LOOKUP_HEADERS).
|
||||||
|
LookupSource string
|
||||||
|
LookupDBPath string
|
||||||
|
LookupTimeout time.Duration
|
||||||
|
AddLookupHeaders bool
|
||||||
// DeniedCountries are the countries whose clients are refused
|
// DeniedCountries are the countries whose clients are refused
|
||||||
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
|
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
|
||||||
// empty, are the only countries whose clients are let through
|
// empty, are the only countries whose clients are let through
|
||||||
@@ -95,6 +120,22 @@ type Config struct {
|
|||||||
// capitals, as GeoJS gives them.
|
// capitals, as GeoJS gives them.
|
||||||
DeniedCountries []string
|
DeniedCountries []string
|
||||||
ExclusivelyAllowedCountries []string
|
ExclusivelyAllowedCountries []string
|
||||||
|
// The biased thresholds. ASNLimitPercent and CountryLimitPercent give
|
||||||
|
// the clients of the AS numbers and the countries they list that
|
||||||
|
// percentage of every rate limit and byte limit
|
||||||
|
// (SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT).
|
||||||
|
// ASNBytesPercent and CountryBytesPercent give those they list a
|
||||||
|
// percentage of the byte limits in place of that one
|
||||||
|
// (SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT). Each holds
|
||||||
|
// percentages from 0 to 100, by AS number, written as AS64496, or by
|
||||||
|
// country, a two-letter code in capitals, as the lookup gives them.
|
||||||
|
// UnknownLimitPercent is the percentage of every limit a client without
|
||||||
|
// a country gets (SWWAF_UNKNOWN_LIMIT_PERCENT).
|
||||||
|
ASNLimitPercent map[string]int64
|
||||||
|
CountryLimitPercent map[string]int64
|
||||||
|
ASNBytesPercent map[string]int64
|
||||||
|
CountryBytesPercent map[string]int64
|
||||||
|
UnknownLimitPercent int64
|
||||||
// BanResponse is the status a refused client is answered with, 403
|
// BanResponse is the status a refused client is answered with, 403
|
||||||
// or 429, or 0 to close the connection without an answer
|
// or 429, or 0 to close the connection without an answer
|
||||||
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
|
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
|
||||||
@@ -135,8 +176,8 @@ type Config struct {
|
|||||||
AdminToken string
|
AdminToken string
|
||||||
// MetricsToken is the bearer token a scraper sends for the metrics
|
// MetricsToken is the bearer token a scraper sends for the metrics
|
||||||
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
|
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
|
||||||
// MetricsTopN is how many countries get series of their own in the
|
// MetricsTopN is how many AS numbers and how many countries get series
|
||||||
// metrics (SWWAF_METRICS_TOP_N).
|
// of their own in the metrics (SWWAF_METRICS_TOP_N).
|
||||||
MetricsToken string
|
MetricsToken string
|
||||||
MetricsTopN int
|
MetricsTopN int
|
||||||
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
|
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
|
||||||
@@ -158,6 +199,27 @@ type Config struct {
|
|||||||
LogRemoteBuffer int
|
LogRemoteBuffer int
|
||||||
LogRemoteFacility int
|
LogRemoteFacility int
|
||||||
LogRemoteAppName string
|
LogRemoteAppName string
|
||||||
|
// AlertWebhookURL is where each alert is posted as JSON
|
||||||
|
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. AlertWebhookHeaders
|
||||||
|
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
|
||||||
|
// AlertSlackWebhookURL is the Slack incoming webhook each alert is
|
||||||
|
// posted to as a message (SWWAF_ALERT_SLACK_WEBHOOK_URL), and
|
||||||
|
// AlertNtfyURL the ntfy topic each is published to
|
||||||
|
// (SWWAF_ALERT_NTFY_URL), each nil while it is unset; AlertNtfyToken,
|
||||||
|
// unless empty, is sent to ntfy with each (SWWAF_ALERT_NTFY_TOKEN).
|
||||||
|
// With none of the three URLs set, no alert is sent. AlertEvents are
|
||||||
|
// the events alerts are sent for (SWWAF_ALERT_EVENTS). A repeat of an
|
||||||
|
// alert within AlertCooldown is held back (SWWAF_ALERT_COOLDOWN), and
|
||||||
|
// so is an alert past AlertMaxPerHour in an hour, for the hour's
|
||||||
|
// summary (SWWAF_ALERT_MAX_PER_HOUR); 0 is off for both.
|
||||||
|
AlertWebhookURL *url.URL
|
||||||
|
AlertWebhookHeaders http.Header
|
||||||
|
AlertSlackWebhookURL *url.URL
|
||||||
|
AlertNtfyURL *url.URL
|
||||||
|
AlertNtfyToken string
|
||||||
|
AlertEvents []string
|
||||||
|
AlertCooldown time.Duration
|
||||||
|
AlertMaxPerHour int
|
||||||
|
|
||||||
// settings are the values read, as given or by default, and the
|
// settings are the values read, as given or by default, and the
|
||||||
// files they were read from, for the log line at start.
|
// files they were read from, for the log line at start.
|
||||||
@@ -168,6 +230,10 @@ type Config struct {
|
|||||||
// off.
|
// off.
|
||||||
const off = "off"
|
const off = "off"
|
||||||
|
|
||||||
|
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
|
||||||
|
// lookup database, the file SWWAF_LOOKUP_DB_PATH names.
|
||||||
|
const fileSource = "file"
|
||||||
|
|
||||||
const (
|
const (
|
||||||
day = 24 * time.Hour
|
day = 24 * time.Hour
|
||||||
kibibyte = 1 << 10
|
kibibyte = 1 << 10
|
||||||
@@ -176,7 +242,8 @@ const (
|
|||||||
ipv4Bits = 32
|
ipv4Bits = 32
|
||||||
// minTokenLength is the fewest characters a token may have.
|
// minTokenLength is the fewest characters a token may have.
|
||||||
minTokenLength = 32
|
minTokenLength = 32
|
||||||
// masked is what the log shows for a token that is set.
|
// masked is what the log shows for a token that is set, and in place of
|
||||||
|
// a secret in another setting.
|
||||||
masked = "********"
|
masked = "********"
|
||||||
// defaultListenAddr and defaultUpstreamURL are the defaults of
|
// defaultListenAddr and defaultUpstreamURL are the defaults of
|
||||||
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
|
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
|
||||||
@@ -208,6 +275,10 @@ var (
|
|||||||
"is taken out of every request by Go's HTTP server, so it can never " +
|
"is taken out of every request by Go's HTTP server, so it can never " +
|
||||||
"be logged")
|
"be logged")
|
||||||
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
|
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
|
||||||
|
errNotLookupSource = errors.New("is not geojs, file or off")
|
||||||
|
errNeedsLookups = errors.New("it needs each client looked up")
|
||||||
|
errNeedsDBPath = errors.New("it names the file to look clients up in")
|
||||||
|
errDBPathUnused = errors.New("only file reads it")
|
||||||
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
|
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
|
||||||
errNotDurationAboveZero = errors.New(
|
errNotDurationAboveZero = errors.New(
|
||||||
"is not a duration above zero, such as 1h or 7d")
|
"is not a duration above zero, such as 1h or 7d")
|
||||||
@@ -220,6 +291,7 @@ var (
|
|||||||
"is not an absolute path, such as /var/lib/smallwebwaf")
|
"is not an absolute path, such as /var/lib/smallwebwaf")
|
||||||
errShortToken = errors.New("is shorter than 32 characters")
|
errShortToken = errors.New("is shorter than 32 characters")
|
||||||
errNotMode = errors.New("is not enforce or observe")
|
errNotMode = errors.New("is not enforce or observe")
|
||||||
|
errNotBytesCount = errors.New("is not response, request or both")
|
||||||
errNotPathPrefix = errors.New(
|
errNotPathPrefix = errors.New(
|
||||||
"is not a path prefix starting with /, such as /assets/")
|
"is not a path prefix starting with /, such as /assets/")
|
||||||
errNotBoolean = errors.New("is not true or false")
|
errNotBoolean = errors.New("is not true or false")
|
||||||
@@ -230,7 +302,25 @@ var (
|
|||||||
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
|
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
|
||||||
errNotAppName = errors.New(
|
errNotAppName = errors.New(
|
||||||
"is not 1 to 48 printable ASCII characters without a space, such as gitea")
|
"is not 1 to 48 printable ASCII characters without a space, such as gitea")
|
||||||
errSetTwice = errors.New("set only one of them")
|
errSetTwice = errors.New("set only one of them")
|
||||||
|
errNotWebhookURL = errors.New(
|
||||||
|
"is not an http or https URL without a user or a fragment, " +
|
||||||
|
"such as https://alerts.example/smallwebwaf")
|
||||||
|
errNotWebhookHeader = errors.New(
|
||||||
|
"is not a header name followed by : and the header's value, " +
|
||||||
|
"such as Authorization:Bearer <token>")
|
||||||
|
errControlCharacter = errors.New(
|
||||||
|
"holds a control character, such as the carriage return of a Windows line end")
|
||||||
|
errNotAlertEvent = errors.New(
|
||||||
|
"is not ban, permanent_ban, waf_block, anomaly, reputation_hit, " +
|
||||||
|
"source_failure or file_error")
|
||||||
|
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
|
||||||
|
errNotUTF8 = errors.New("is not valid UTF-8")
|
||||||
|
errNotASN = errors.New("is not an AS number such as AS64496")
|
||||||
|
errNotPercentItem = errors.New(
|
||||||
|
"is not a code, : and a percentage, such as AS64496:50 or cn:25")
|
||||||
|
errNotPercent = errors.New("is not a percentage, a whole number from 0 to 100")
|
||||||
|
errListedTwice = errors.New("is listed twice")
|
||||||
)
|
)
|
||||||
|
|
||||||
// FromEnvironment reads the settings with lookupEnv, normally
|
// FromEnvironment reads the settings with lookupEnv, normally
|
||||||
@@ -238,13 +328,14 @@ var (
|
|||||||
// named by the setting's name with _FILE added names the file, which is
|
// named by the setting's name with _FILE added names the file, which is
|
||||||
// read now (see lookup). A setting that is not set takes its default. A
|
// read now (see lookup). A setting that is not set takes its default. A
|
||||||
// setting that is set but invalid is an error that names it.
|
// setting that is set but invalid is an error that names it.
|
||||||
|
//
|
||||||
|
//nolint:funlen // one line for each setting, a list that grows with them
|
||||||
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||||
env := &environment{lookupEnv: lookupEnv}
|
env := &environment{lookupEnv: lookupEnv}
|
||||||
hostname, _ := os.Hostname() // "" when the host has no name to give
|
|
||||||
cfg := &Config{
|
cfg := &Config{
|
||||||
ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
|
ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
|
||||||
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
|
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
|
||||||
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
|
InstanceName: env.instanceName(),
|
||||||
Observe: env.observe("SWWAF_MODE", "enforce"),
|
Observe: env.observe("SWWAF_MODE", "enforce"),
|
||||||
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
|
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
|
||||||
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
|
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
|
||||||
@@ -263,9 +354,22 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
|||||||
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
|
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
|
||||||
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
|
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
|
||||||
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
|
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
|
||||||
|
BytesLimitPerMinute: env.size("SWWAF_BYTES_LIMIT_PER_MINUTE", "10G"),
|
||||||
|
BytesLimitPerHour: env.size("SWWAF_BYTES_LIMIT_PER_HOUR", "20G"),
|
||||||
|
BytesLimitPerDay: env.size("SWWAF_BYTES_LIMIT_PER_DAY", "50G"),
|
||||||
|
BytesCount: env.bytesCount("SWWAF_BYTES_COUNT", "both"),
|
||||||
|
LookupSource: env.lookupSource("SWWAF_LOOKUP_SOURCE", "geojs"),
|
||||||
|
LookupDBPath: env.value("SWWAF_LOOKUP_DB_PATH", ""),
|
||||||
|
LookupTimeout: env.durationNotOff("SWWAF_LOOKUP_TIMEOUT", "1s"),
|
||||||
|
AddLookupHeaders: env.boolean("SWWAF_ADD_LOOKUP_HEADERS", "false"),
|
||||||
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
|
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
|
||||||
ExclusivelyAllowedCountries: env.countries(
|
ExclusivelyAllowedCountries: env.countries(
|
||||||
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
|
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
|
||||||
|
ASNLimitPercent: env.percents("SWWAF_ASN_LIMIT_PERCENT", parseASN),
|
||||||
|
CountryLimitPercent: env.percents("SWWAF_COUNTRY_LIMIT_PERCENT", parseCountry),
|
||||||
|
ASNBytesPercent: env.percents("SWWAF_ASN_BYTES_PERCENT", parseASN),
|
||||||
|
CountryBytesPercent: env.percents("SWWAF_COUNTRY_BYTES_PERCENT", parseCountry),
|
||||||
|
UnknownLimitPercent: env.percent("SWWAF_UNKNOWN_LIMIT_PERCENT", "100"),
|
||||||
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
|
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
|
||||||
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
|
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
|
||||||
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
|
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
|
||||||
@@ -278,26 +382,32 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
|||||||
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
|
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
|
||||||
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
|
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
|
||||||
"accept,accept-language,accept-encoding,content-type,origin,range"),
|
"accept,accept-language,accept-encoding,content-type,origin,range"),
|
||||||
AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
|
AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
|
||||||
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
|
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
|
||||||
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
|
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
|
||||||
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
|
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
|
||||||
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
|
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
|
||||||
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
|
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
|
||||||
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
|
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
|
||||||
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
|
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
|
||||||
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
|
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
|
||||||
|
AlertWebhookURL: env.webhookURL("SWWAF_ALERT_WEBHOOK_URL"),
|
||||||
|
AlertWebhookHeaders: env.webhookHeaders("SWWAF_ALERT_WEBHOOK_HEADERS"),
|
||||||
|
AlertSlackWebhookURL: env.webhookURL("SWWAF_ALERT_SLACK_WEBHOOK_URL"),
|
||||||
|
AlertNtfyURL: env.webhookURL("SWWAF_ALERT_NTFY_URL"),
|
||||||
|
AlertNtfyToken: env.secret("SWWAF_ALERT_NTFY_TOKEN"),
|
||||||
|
AlertEvents: env.alertEvents("SWWAF_ALERT_EVENTS",
|
||||||
|
strings.Join(alerts.Events(), ",")),
|
||||||
|
AlertCooldown: env.duration("SWWAF_ALERT_COOLDOWN", "15m"),
|
||||||
|
AlertMaxPerHour: env.numberOrOff("SWWAF_ALERT_MAX_PER_HOUR", "60"),
|
||||||
}
|
}
|
||||||
|
|
||||||
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
|
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
|
||||||
cfg.InstanceName, cfg.LogRemoteURL != nil)
|
cfg.InstanceName, cfg.LogRemoteURL != nil)
|
||||||
|
|
||||||
for _, country := range cfg.ExclusivelyAllowedCountries {
|
env.checkInstanceNameForNtfy(cfg.InstanceName, cfg.AlertNtfyURL != nil)
|
||||||
if slices.Contains(cfg.DeniedCountries, country) {
|
env.checkLookupDBPath(cfg)
|
||||||
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
|
env.checkCountriesAndLookups(cfg)
|
||||||
fmt.Errorf("%q %w", country, errOnBothLists))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if env.err != nil {
|
if env.err != nil {
|
||||||
return nil, env.err
|
return nil, env.err
|
||||||
@@ -326,6 +436,16 @@ func ListenAddrAndUpstreamURL(
|
|||||||
return listenAddr, upstreamURL, nil
|
return listenAddr, upstreamURL, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// InstanceName reads only SWWAF_INSTANCE_NAME, which may be given as a
|
||||||
|
// file, as FromEnvironment does, so that the line saying a setting is
|
||||||
|
// invalid carries it too. A file that cannot be read gives the default
|
||||||
|
// here, and FromEnvironment then stops the start over it.
|
||||||
|
func InstanceName(lookupEnv func(string) (string, bool)) string {
|
||||||
|
env := &environment{lookupEnv: lookupEnv}
|
||||||
|
|
||||||
|
return env.instanceName()
|
||||||
|
}
|
||||||
|
|
||||||
// privateRanges are the private address ranges, the default trusted
|
// privateRanges are the private address ranges, the default trusted
|
||||||
// proxies.
|
// proxies.
|
||||||
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
|
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
|
||||||
@@ -473,6 +593,17 @@ func (e *environment) count(name, defaultValue string) int64 {
|
|||||||
return count
|
return count
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// bytesCount reads the setting that is which body bytes count toward the
|
||||||
|
// byte limits: response, request or both.
|
||||||
|
func (e *environment) bytesCount(name, defaultValue string) string {
|
||||||
|
value := e.value(name, defaultValue)
|
||||||
|
if value != "response" && value != "request" && value != "both" {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", value, errNotBytesCount))
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
// pathPrefixes reads a setting that is a list of path prefixes.
|
// pathPrefixes reads a setting that is a list of path prefixes.
|
||||||
func (e *environment) pathPrefixes(name, defaultValue string) []string {
|
func (e *environment) pathPrefixes(name, defaultValue string) []string {
|
||||||
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
|
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
|
||||||
@@ -489,6 +620,87 @@ func (e *environment) countries(name, defaultValue string) []string {
|
|||||||
return countries
|
return countries
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// percents reads a setting that is a list of AS numbers or countries,
|
||||||
|
// which parseCode reads, each with a percentage. It is empty by default.
|
||||||
|
func (e *environment) percents(
|
||||||
|
name string, parseCode func(string) (string, error),
|
||||||
|
) map[string]int64 {
|
||||||
|
percents, err := parsePercents(e.value(name, ""), parseCode)
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return percents
|
||||||
|
}
|
||||||
|
|
||||||
|
// percent reads a setting that is a percentage, from 0 to 100.
|
||||||
|
func (e *environment) percent(name, defaultValue string) int64 {
|
||||||
|
percent, err := parsePercent(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return percent
|
||||||
|
}
|
||||||
|
|
||||||
|
// lookupSource reads the setting that is where clients are looked up:
|
||||||
|
// geojs, file, or off.
|
||||||
|
func (e *environment) lookupSource(name, defaultValue string) string {
|
||||||
|
source := e.value(name, defaultValue)
|
||||||
|
if source != "geojs" && source != fileSource && source != off {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", source, errNotLookupSource))
|
||||||
|
}
|
||||||
|
|
||||||
|
return source
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkLookupDBPath refuses SWWAF_LOOKUP_SOURCE=file without
|
||||||
|
// SWWAF_LOOKUP_DB_PATH, and SWWAF_LOOKUP_DB_PATH with any other source:
|
||||||
|
// one source at a time.
|
||||||
|
func (e *environment) checkLookupDBPath(cfg *Config) {
|
||||||
|
switch {
|
||||||
|
case cfg.LookupSource == fileSource && cfg.LookupDBPath == "":
|
||||||
|
e.check("SWWAF_LOOKUP_SOURCE", fmt.Errorf(
|
||||||
|
"is file while SWWAF_LOOKUP_DB_PATH is unset; %w", errNeedsDBPath))
|
||||||
|
case cfg.LookupSource != fileSource && cfg.LookupDBPath != "":
|
||||||
|
e.check("SWWAF_LOOKUP_DB_PATH", fmt.Errorf(
|
||||||
|
"is set while SWWAF_LOOKUP_SOURCE is %s; %w", cfg.LookupSource, errDBPathUnused))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkCountriesAndLookups refuses a country on both country lists, and,
|
||||||
|
// while SWWAF_LOOKUP_SOURCE is off, each setting that needs clients looked
|
||||||
|
// up: the country lists, SWWAF_ADD_LOOKUP_HEADERS, and the biased
|
||||||
|
// thresholds, of which SWWAF_UNKNOWN_LIMIT_PERCENT needs them only below
|
||||||
|
// 100, where it lowers a limit.
|
||||||
|
func (e *environment) checkCountriesAndLookups(cfg *Config) {
|
||||||
|
for _, country := range cfg.ExclusivelyAllowedCountries {
|
||||||
|
if slices.Contains(cfg.DeniedCountries, country) {
|
||||||
|
e.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
|
||||||
|
fmt.Errorf("%q %w", country, errOnBothLists))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.LookupSource != off {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, setting := range []struct {
|
||||||
|
name string
|
||||||
|
set bool
|
||||||
|
}{
|
||||||
|
{"SWWAF_DENIED_COUNTRIES", len(cfg.DeniedCountries) > 0},
|
||||||
|
{"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", len(cfg.ExclusivelyAllowedCountries) > 0},
|
||||||
|
{"SWWAF_ADD_LOOKUP_HEADERS", cfg.AddLookupHeaders},
|
||||||
|
{"SWWAF_ASN_LIMIT_PERCENT", len(cfg.ASNLimitPercent) > 0},
|
||||||
|
{"SWWAF_COUNTRY_LIMIT_PERCENT", len(cfg.CountryLimitPercent) > 0},
|
||||||
|
{"SWWAF_ASN_BYTES_PERCENT", len(cfg.ASNBytesPercent) > 0},
|
||||||
|
{"SWWAF_COUNTRY_BYTES_PERCENT", len(cfg.CountryBytesPercent) > 0},
|
||||||
|
{"SWWAF_UNKNOWN_LIMIT_PERCENT", cfg.UnknownLimitPercent < 100},
|
||||||
|
} {
|
||||||
|
if setting.set {
|
||||||
|
e.check(setting.name, fmt.Errorf("is set while SWWAF_LOOKUP_SOURCE is off; %w",
|
||||||
|
errNeedsLookups))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// headerNames reads a setting that is a list of header names, and
|
// headerNames reads a setting that is a list of header names, and
|
||||||
// returns them in lower case.
|
// returns them in lower case.
|
||||||
func (e *environment) headerNames(name, defaultValue string) []string {
|
func (e *environment) headerNames(name, defaultValue string) []string {
|
||||||
@@ -613,6 +825,19 @@ func (e *environment) facility(name, defaultValue string) int {
|
|||||||
return number
|
return number
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// instanceName reads SWWAF_INSTANCE_NAME, by default the host's name. It
|
||||||
|
// must be valid UTF-8: the metrics library panics on a label that is not.
|
||||||
|
func (e *environment) instanceName() string {
|
||||||
|
hostname, _ := os.Hostname() // "" when the host has no name to give
|
||||||
|
|
||||||
|
value := e.value("SWWAF_INSTANCE_NAME", hostname)
|
||||||
|
if !utf8.ValidString(value) {
|
||||||
|
e.check("SWWAF_INSTANCE_NAME", fmt.Errorf("%q %w", value, errNotUTF8))
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
// appName reads the setting that is the APP-NAME of the records the log
|
// appName reads the setting that is the APP-NAME of the records the log
|
||||||
// lines are sent in, by default the instance name. Its value is checked
|
// lines are sent in, by default the instance name. Its value is checked
|
||||||
// when it is set, and, while lines are sent, when it is the instance name.
|
// when it is set, and, while lines are sent, when it is the instance name.
|
||||||
@@ -636,6 +861,83 @@ func (e *environment) appName(name, instanceName string, sending bool) string {
|
|||||||
return value
|
return value
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// checkInstanceNameForNtfy refuses an instance name that holds a control
|
||||||
|
// character while ntfySet, SWWAF_ALERT_NTFY_URL being set: ntfy is sent
|
||||||
|
// the instance name in a header, which cannot hold one.
|
||||||
|
func (e *environment) checkInstanceNameForNtfy(instanceName string, ntfySet bool) {
|
||||||
|
if ntfySet && strings.ContainsFunc(instanceName, unicode.IsControl) {
|
||||||
|
e.check("SWWAF_INSTANCE_NAME", fmt.Errorf(
|
||||||
|
"%q %w, and is sent to ntfy in a header while SWWAF_ALERT_NTFY_URL is set",
|
||||||
|
instanceName, errControlCharacter))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// webhookURL reads a setting that is a URL each alert is posted to:
|
||||||
|
// SWWAF_ALERT_WEBHOOK_URL, SWWAF_ALERT_SLACK_WEBHOOK_URL or
|
||||||
|
// SWWAF_ALERT_NTFY_URL. Unset or empty, it is nil, and no alert is posted
|
||||||
|
// there. The log shows ******** in place of its path and query, and an
|
||||||
|
// error shows none of it, since a webhook or an ntfy topic can carry its
|
||||||
|
// secret there.
|
||||||
|
func (e *environment) webhookURL(name string) *url.URL {
|
||||||
|
value, _ := e.lookup(name)
|
||||||
|
webhook, logged, err := parseWebhookURL(value)
|
||||||
|
e.settings = append(e.settings, slog.String(name, logged))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return webhook
|
||||||
|
}
|
||||||
|
|
||||||
|
// webhookHeaders reads the setting that is the headers sent with each
|
||||||
|
// alert. The log shows each header's value as ********, since a header
|
||||||
|
// such as Authorization carries a secret.
|
||||||
|
func (e *environment) webhookHeaders(name string) http.Header {
|
||||||
|
value, _ := e.lookup(name)
|
||||||
|
headers, logged, err := parseWebhookHeaders(value)
|
||||||
|
e.settings = append(e.settings, slog.String(name, logged))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
|
// secret reads a setting that is a secret another service gave, such as
|
||||||
|
// an ntfy token, "" while it is unset. It is sent in a header, which
|
||||||
|
// cannot hold a control character, so one in it is an error. The log
|
||||||
|
// shows ******** in place of a value that is not empty, and an error
|
||||||
|
// shows none of it.
|
||||||
|
func (e *environment) secret(name string) string {
|
||||||
|
value, _ := e.lookup(name)
|
||||||
|
|
||||||
|
logged := ""
|
||||||
|
if value != "" {
|
||||||
|
logged = masked
|
||||||
|
}
|
||||||
|
|
||||||
|
e.settings = append(e.settings, slog.String(name, logged))
|
||||||
|
|
||||||
|
if strings.ContainsFunc(value, unicode.IsControl) {
|
||||||
|
e.check(name, errControlCharacter)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// alertEvents reads the setting that is the events alerts are sent for.
|
||||||
|
func (e *environment) alertEvents(name, defaultValue string) []string {
|
||||||
|
events, err := parseAlertEvents(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return events
|
||||||
|
}
|
||||||
|
|
||||||
|
// numberOrOff reads a setting that is a whole number above zero, or off,
|
||||||
|
// which is 0.
|
||||||
|
func (e *environment) numberOrOff(name, defaultValue string) int {
|
||||||
|
number, err := parseNumberOrOff(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return number
|
||||||
|
}
|
||||||
|
|
||||||
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
||||||
// whole number of days such as 7d, or off.
|
// whole number of days such as 7d, or off.
|
||||||
func parseDuration(value string) (time.Duration, error) {
|
func parseDuration(value string) (time.Duration, error) {
|
||||||
@@ -900,13 +1202,12 @@ func parseCountries(value string) ([]string, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
known := strings.Fields(countryCodes)
|
|
||||||
countries := make([]string, 0, len(items))
|
countries := make([]string, 0, len(items))
|
||||||
|
|
||||||
for _, item := range items {
|
for _, item := range items {
|
||||||
country := strings.ToUpper(item)
|
country, err := parseCountry(item)
|
||||||
if !slices.Contains(known, country) {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("%q %w", item, errNotCountry)
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
countries = append(countries, country)
|
countries = append(countries, country)
|
||||||
@@ -915,6 +1216,81 @@ func parseCountries(value string) ([]string, error) {
|
|||||||
return countries, nil
|
return countries, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseCountry reads a country code in either case, and returns it in
|
||||||
|
// capitals.
|
||||||
|
func parseCountry(value string) (string, error) {
|
||||||
|
country := strings.ToUpper(value)
|
||||||
|
if !slices.Contains(strings.Fields(countryCodes), country) {
|
||||||
|
return "", fmt.Errorf("%q %w", value, errNotCountry)
|
||||||
|
}
|
||||||
|
|
||||||
|
return country, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseASN reads an AS number such as AS64496, in either case, and
|
||||||
|
// returns it as the lookup gives it: AS and the number, in capitals and
|
||||||
|
// without leading zeros.
|
||||||
|
func parseASN(value string) (string, error) {
|
||||||
|
digits, hasAS := strings.CutPrefix(strings.ToUpper(value), "AS")
|
||||||
|
|
||||||
|
number, err := strconv.ParseUint(digits, 10, 32)
|
||||||
|
if !hasAS || err != nil {
|
||||||
|
return "", fmt.Errorf("%q %w", value, errNotASN)
|
||||||
|
}
|
||||||
|
|
||||||
|
return "AS" + strconv.FormatUint(number, 10), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsePercents reads a comma-separated list of items, each an AS number
|
||||||
|
// or a country, which parseCode reads, then : and a percentage, such as
|
||||||
|
// AS64496:50 or cn:25, and returns each one's percentage. An empty value
|
||||||
|
// is an empty list. An AS number or country listed twice is an error.
|
||||||
|
func parsePercents(
|
||||||
|
value string, parseCode func(string) (string, error),
|
||||||
|
) (map[string]int64, error) {
|
||||||
|
items, err := parseList(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
percents := make(map[string]int64, len(items))
|
||||||
|
|
||||||
|
for _, item := range items {
|
||||||
|
codeText, percentText, found := strings.Cut(item, ":")
|
||||||
|
if !found {
|
||||||
|
return nil, fmt.Errorf("%q %w", item, errNotPercentItem)
|
||||||
|
}
|
||||||
|
|
||||||
|
code, err := parseCode(codeText)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
percent, err := parsePercent(percentText)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, listed := percents[code]; listed {
|
||||||
|
return nil, fmt.Errorf("%q %w", codeText, errListedTwice)
|
||||||
|
}
|
||||||
|
|
||||||
|
percents[code] = percent
|
||||||
|
}
|
||||||
|
|
||||||
|
return percents, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsePercent reads a percentage, a whole number from 0 to 100.
|
||||||
|
func parsePercent(value string) (int64, error) {
|
||||||
|
percent, err := strconv.ParseInt(value, 10, 64)
|
||||||
|
if err != nil || percent < 0 || percent > 100 {
|
||||||
|
return 0, fmt.Errorf("%q %w", value, errNotPercent)
|
||||||
|
}
|
||||||
|
|
||||||
|
return percent, nil
|
||||||
|
}
|
||||||
|
|
||||||
// headerNameChars are the characters RFC 9110 allows in a header name:
|
// headerNameChars are the characters RFC 9110 allows in a header name:
|
||||||
// letters, digits and these marks.
|
// letters, digits and these marks.
|
||||||
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
|
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
|
||||||
@@ -1051,6 +1427,96 @@ func parseFacility(value string) (int, error) {
|
|||||||
return number, nil
|
return number, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseWebhookURL reads where each alert is posted: http or https, a
|
||||||
|
// host, and an optional port from 1 to 65535, path and query, without a
|
||||||
|
// user or a fragment. It returns the URL, and how the log shows it: its
|
||||||
|
// scheme and host, and ******** in place of its path and query, if it has
|
||||||
|
// either. An error shows no part of the value. An empty value is no URL.
|
||||||
|
func parseWebhookURL(value string) (*url.URL, string, error) {
|
||||||
|
if value == "" {
|
||||||
|
return nil, "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
webhook, err := url.Parse(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, "", errNotWebhookURL
|
||||||
|
}
|
||||||
|
|
||||||
|
port, err := strconv.ParseUint(webhook.Port(), 10, 16)
|
||||||
|
|
||||||
|
valid := (webhook.Scheme == "http" || webhook.Scheme == "https") &&
|
||||||
|
webhook.Hostname() != "" && (webhook.Port() == "" || (err == nil && port != 0)) &&
|
||||||
|
webhook.User == nil && webhook.Opaque == "" && webhook.Fragment == ""
|
||||||
|
if !valid {
|
||||||
|
return nil, "", errNotWebhookURL
|
||||||
|
}
|
||||||
|
|
||||||
|
logged := webhook.Scheme + "://" + webhook.Host
|
||||||
|
if webhook.Path != "" || webhook.RawQuery != "" {
|
||||||
|
logged += "/" + masked
|
||||||
|
}
|
||||||
|
|
||||||
|
return webhook, logged, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseWebhookHeaders reads a comma-separated list of headers, each its
|
||||||
|
// name, :, and its value, and returns them, and how the log shows them,
|
||||||
|
// with each value as ********. An error names the item by its place in
|
||||||
|
// the list, so that it shows no value. An empty value is an empty list.
|
||||||
|
func parseWebhookHeaders(value string) (http.Header, string, error) {
|
||||||
|
headers := http.Header{}
|
||||||
|
if strings.TrimSpace(value) == "" {
|
||||||
|
return headers, "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
logged := []string{}
|
||||||
|
|
||||||
|
for i, item := range strings.Split(value, ",") {
|
||||||
|
name, headerValue, found := strings.Cut(item, ":")
|
||||||
|
name = strings.TrimSpace(name)
|
||||||
|
|
||||||
|
if !found || !IsHeaderName(name) || strings.ContainsAny(headerValue, "\r\n\x00") {
|
||||||
|
return nil, "", fmt.Errorf("item %d %w", i+1, errNotWebhookHeader)
|
||||||
|
}
|
||||||
|
|
||||||
|
headers.Add(name, strings.TrimSpace(headerValue))
|
||||||
|
logged = append(logged, name+":"+masked)
|
||||||
|
}
|
||||||
|
|
||||||
|
return headers, strings.Join(logged, ","), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseAlertEvents reads a comma-separated list of the events alerts can
|
||||||
|
// be sent for.
|
||||||
|
func parseAlertEvents(value string) ([]string, error) {
|
||||||
|
events, err := parseList(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, event := range events {
|
||||||
|
if !slices.Contains(alerts.Events(), event) {
|
||||||
|
return nil, fmt.Errorf("%q %w", event, errNotAlertEvent)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return events, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseNumberOrOff reads a whole number above zero, or off, which is 0.
|
||||||
|
func parseNumberOrOff(value string) (int, error) {
|
||||||
|
if value == off {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
n, err := strconv.Atoi(value)
|
||||||
|
if err != nil || n <= 0 {
|
||||||
|
return 0, fmt.Errorf("%q %w", value, errNotNumberOrOff)
|
||||||
|
}
|
||||||
|
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
// appNameMaxLength is the most characters RFC 5424 allows in an
|
// appNameMaxLength is the most characters RFC 5424 allows in an
|
||||||
// APP-NAME.
|
// APP-NAME.
|
||||||
const appNameMaxLength = 48
|
const appNameMaxLength = 48
|
||||||
|
|||||||
+597
-16
@@ -6,9 +6,11 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"maps"
|
"maps"
|
||||||
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -38,8 +40,21 @@ const (
|
|||||||
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
|
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
|
||||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
||||||
|
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
|
||||||
|
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
|
||||||
|
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
|
||||||
|
bytesCount = "SWWAF_BYTES_COUNT"
|
||||||
|
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
||||||
|
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
||||||
|
lookupTimeout = "SWWAF_LOOKUP_TIMEOUT"
|
||||||
|
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
|
||||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||||
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
|
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
|
||||||
|
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
|
||||||
|
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
|
||||||
|
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
|
||||||
|
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
|
||||||
|
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
|
||||||
banResponse = "SWWAF_BAN_RESPONSE"
|
banResponse = "SWWAF_BAN_RESPONSE"
|
||||||
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
|
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
|
||||||
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
|
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
|
||||||
@@ -62,6 +77,22 @@ const (
|
|||||||
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
|
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
|
||||||
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
|
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
|
||||||
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
|
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
|
||||||
|
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
|
||||||
|
alertWebhookHeaders = "SWWAF_ALERT_WEBHOOK_HEADERS"
|
||||||
|
alertSlackWebhookURL = "SWWAF_ALERT_SLACK_WEBHOOK_URL"
|
||||||
|
alertNtfyURL = "SWWAF_ALERT_NTFY_URL"
|
||||||
|
alertNtfyToken = "SWWAF_ALERT_NTFY_TOKEN" //nolint:gosec // the setting's name
|
||||||
|
alertEvents = "SWWAF_ALERT_EVENTS"
|
||||||
|
alertCooldown = "SWWAF_ALERT_COOLDOWN"
|
||||||
|
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
|
||||||
|
)
|
||||||
|
|
||||||
|
// defaultAlertEvents is the default of SWWAF_ALERT_EVENTS, and
|
||||||
|
// defaultAlertCooldown that of SWWAF_ALERT_COOLDOWN.
|
||||||
|
const (
|
||||||
|
defaultAlertEvents = "ban,permanent_ban,waf_block,anomaly,reputation_hit," +
|
||||||
|
"source_failure,file_error"
|
||||||
|
defaultAlertCooldown = "15m"
|
||||||
)
|
)
|
||||||
|
|
||||||
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
|
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
|
||||||
@@ -99,6 +130,16 @@ const (
|
|||||||
// off switches a timeout, a size limit or a rate limit off.
|
// off switches a timeout, a size limit or a rate limit off.
|
||||||
const off = "off"
|
const off = "off"
|
||||||
|
|
||||||
|
// enabled is true, as a setting's value.
|
||||||
|
const enabled = "true"
|
||||||
|
|
||||||
|
// defaultLookupSource is the default of SWWAF_LOOKUP_SOURCE, and
|
||||||
|
// fileSource the source that is the lookup database.
|
||||||
|
const (
|
||||||
|
defaultLookupSource = "geojs"
|
||||||
|
fileSource = "file"
|
||||||
|
)
|
||||||
|
|
||||||
// environment is a set of environment variables, for FromEnvironment.
|
// environment is a set of environment variables, for FromEnvironment.
|
||||||
type environment map[string]string
|
type environment map[string]string
|
||||||
|
|
||||||
@@ -155,6 +196,13 @@ func TestDefaults(t *testing.T) {
|
|||||||
RulesDir: "/etc/smallwebwaf/rules.d",
|
RulesDir: "/etc/smallwebwaf/rules.d",
|
||||||
RulesEnabled: true,
|
RulesEnabled: true,
|
||||||
})
|
})
|
||||||
|
wantLookupSettings(t, cfg, config.Config{
|
||||||
|
LookupSource: defaultLookupSource, LookupTimeout: time.Second,
|
||||||
|
})
|
||||||
|
wantByteLimitSettings(t, cfg, config.Config{
|
||||||
|
BytesLimitPerMinute: 10 << 30, BytesLimitPerHour: 20 << 30,
|
||||||
|
BytesLimitPerDay: 50 << 30, BytesCount: "both",
|
||||||
|
})
|
||||||
|
|
||||||
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
||||||
t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL)
|
t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL)
|
||||||
@@ -185,6 +233,22 @@ func TestDefaults(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestNoSingleRequestBreaksAByteLimitAtTheDefaults(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{})
|
||||||
|
|
||||||
|
// The largest request body and the largest response, both counted.
|
||||||
|
largest := cfg.RequestMaxBytes + cfg.ResponseMaxBytes
|
||||||
|
for _, limit := range []int64{
|
||||||
|
cfg.BytesLimitPerMinute, cfg.BytesLimitPerHour, cfg.BytesLimitPerDay,
|
||||||
|
} {
|
||||||
|
if largest > limit {
|
||||||
|
t.Errorf("a request of %d bytes breaks the byte limit of %d", largest, limit)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestValuesAsSet(t *testing.T) {
|
func TestValuesAsSet(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -267,6 +331,22 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
|
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestByteLimitSettingsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
bytesLimitPerMinute: "512M",
|
||||||
|
bytesLimitPerHour: off,
|
||||||
|
bytesLimitPerDay: "100000",
|
||||||
|
bytesCount: "response",
|
||||||
|
})
|
||||||
|
|
||||||
|
wantByteLimitSettings(t, cfg, config.Config{
|
||||||
|
BytesLimitPerMinute: 512 << 20, BytesLimitPerHour: 0,
|
||||||
|
BytesLimitPerDay: 100000, BytesCount: "response",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestRateLimitExemptPathsAsSet(t *testing.T) {
|
func TestRateLimitExemptPathsAsSet(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -477,6 +557,267 @@ func TestAppNameSetStopsTheStartWhileSending(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAlertSettingsDefaults(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{})
|
||||||
|
|
||||||
|
if cfg.AlertWebhookURL != nil || len(cfg.AlertWebhookHeaders) != 0 ||
|
||||||
|
cfg.AlertSlackWebhookURL != nil || cfg.AlertNtfyURL != nil ||
|
||||||
|
cfg.AlertNtfyToken != "" ||
|
||||||
|
strings.Join(cfg.AlertEvents, ",") != defaultAlertEvents ||
|
||||||
|
cfg.AlertCooldown != 15*time.Minute || cfg.AlertMaxPerHour != 60 {
|
||||||
|
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d, want no URLs, "+
|
||||||
|
"no headers, no token, %s, 15m and 60", cfg.AlertWebhookURL,
|
||||||
|
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
|
||||||
|
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour,
|
||||||
|
defaultAlertEvents)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAlertSettingsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
webhook = "https://alerts.example:8443/hooks/waf?team=ops"
|
||||||
|
slack = "https://hooks.slack.example/services/T0123/B4567/abcdef"
|
||||||
|
ntfy = "https://ntfy.example/smallwebwaf-alerts"
|
||||||
|
token = "tk_0123456789abcdefghijklmnopq"
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
alertWebhookURL: webhook,
|
||||||
|
alertWebhookHeaders: "Authorization: Bearer abc:def , x-team:ops",
|
||||||
|
alertSlackWebhookURL: slack,
|
||||||
|
alertNtfyURL: ntfy,
|
||||||
|
alertNtfyToken: token,
|
||||||
|
alertEvents: "ban, file_error",
|
||||||
|
alertCooldown: "1h",
|
||||||
|
alertMaxPerHour: "10",
|
||||||
|
})
|
||||||
|
|
||||||
|
headers := http.Header{"Authorization": {"Bearer abc:def"}, "X-Team": {"ops"}}
|
||||||
|
if cfg.AlertWebhookURL.String() != webhook ||
|
||||||
|
!reflect.DeepEqual(cfg.AlertWebhookHeaders, headers) ||
|
||||||
|
cfg.AlertSlackWebhookURL.String() != slack || cfg.AlertNtfyURL.String() != ntfy ||
|
||||||
|
cfg.AlertNtfyToken != token ||
|
||||||
|
!slices.Equal(cfg.AlertEvents, []string{"ban", "file_error"}) ||
|
||||||
|
cfg.AlertCooldown != time.Hour || cfg.AlertMaxPerHour != 10 {
|
||||||
|
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d", cfg.AlertWebhookURL,
|
||||||
|
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
|
||||||
|
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAlertSettingsSetEmptyOrOff(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
alertWebhookURL: "", alertSlackWebhookURL: "", alertNtfyURL: "",
|
||||||
|
alertNtfyToken: "", alertEvents: "", alertCooldown: off, alertMaxPerHour: off,
|
||||||
|
})
|
||||||
|
if cfg.AlertWebhookURL != nil || cfg.AlertSlackWebhookURL != nil ||
|
||||||
|
cfg.AlertNtfyURL != nil || cfg.AlertNtfyToken != "" ||
|
||||||
|
len(cfg.AlertEvents) != 0 || cfg.AlertCooldown != 0 || cfg.AlertMaxPerHour != 0 {
|
||||||
|
t.Errorf("set empty or off, alert settings %v, %v, %v, %q, %v, %s and %d",
|
||||||
|
cfg.AlertWebhookURL, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
|
||||||
|
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInvalidAlertSettingStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantStartStopped(t, []struct{ name, value string }{
|
||||||
|
{alertWebhookURL, "alerts.example/smallwebwaf"},
|
||||||
|
{alertWebhookURL, "ftp://alerts.example/"},
|
||||||
|
{alertWebhookURL, "https:///smallwebwaf"},
|
||||||
|
{alertWebhookURL, "https://user:password@alerts.example/"},
|
||||||
|
{alertWebhookURL, "https://alerts.example/#top"},
|
||||||
|
{alertWebhookURL, "https://alerts.example:0/"},
|
||||||
|
{alertWebhookURL, "https://alerts.example:65536/"},
|
||||||
|
{alertSlackWebhookURL, "hooks.slack.example/services/T0123"},
|
||||||
|
{alertSlackWebhookURL, "https://user:password@hooks.slack.example/"},
|
||||||
|
{alertNtfyURL, "ntfy://ntfy.example/smallwebwaf-alerts"},
|
||||||
|
{alertNtfyURL, "https://ntfy.example/smallwebwaf-alerts#top"},
|
||||||
|
{alertWebhookHeaders, "Authorization"},
|
||||||
|
{alertWebhookHeaders, "X Team:ops"},
|
||||||
|
{alertWebhookHeaders, ":ops"},
|
||||||
|
{alertWebhookHeaders, "X-Team:ops,"},
|
||||||
|
{alertWebhookHeaders, "X-Team:o\r\nps"},
|
||||||
|
{alertEvents, "bans"},
|
||||||
|
{alertEvents, "summary"},
|
||||||
|
{alertEvents, "ban,,file_error"},
|
||||||
|
{alertCooldown, "0"},
|
||||||
|
{alertCooldown, "soon"},
|
||||||
|
{alertMaxPerHour, "0"},
|
||||||
|
{alertMaxPerHour, "-1"},
|
||||||
|
{alertMaxPerHour, "1.5"},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebhookHeadersAreLoggedMaskedAndNeverShown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const secret = "Bearer 0123456789abcdef"
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
alertWebhookHeaders: "Authorization:" + secret + ",X-Team:ops",
|
||||||
|
})
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
logged := out.String()
|
||||||
|
if strings.Contains(logged, secret) || strings.Contains(logged, "ops") ||
|
||||||
|
!strings.Contains(logged,
|
||||||
|
`"`+alertWebhookHeaders+`":"Authorization:********,X-Team:********"`) {
|
||||||
|
t.Errorf("the headers are not logged masked: %s", logged)
|
||||||
|
}
|
||||||
|
|
||||||
|
// An item that is not a header is named by its place, not shown.
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
alertWebhookHeaders: "X-Team:ops," + secret,
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want := alertWebhookHeaders + ": item 2 is not a header name followed by : " +
|
||||||
|
"and the header's value, such as Authorization:Bearer <token>"
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebhookURLIsLoggedWithoutItsPathOrQueryAndNeverShown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const secret = "T0123/B4567/abcdef"
|
||||||
|
|
||||||
|
for value, want := range map[string]string{
|
||||||
|
"https://hooks.example/services/" + secret: "https://hooks.example/********",
|
||||||
|
"https://hooks.example:8443?token=" + secret: "https://hooks.example:8443/********",
|
||||||
|
"http://[2001:db8::1]:8080": "http://[2001:db8::1]:8080",
|
||||||
|
} {
|
||||||
|
cfg := fromEnvironment(t, environment{alertWebhookURL: value})
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
logged := out.String()
|
||||||
|
if strings.Contains(logged, secret) ||
|
||||||
|
!strings.Contains(logged, `"`+alertWebhookURL+`":"`+want+`"`) {
|
||||||
|
t.Errorf("%s is not logged as %s: %s", value, want, logged)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A value that is not such a URL is not shown either.
|
||||||
|
for _, value := range []string{
|
||||||
|
"ftp://hooks.example/services/" + secret,
|
||||||
|
"https://hooks.example/services/%zz" + secret,
|
||||||
|
} {
|
||||||
|
_, err := config.FromEnvironment(environment{alertWebhookURL: value}.lookupEnv)
|
||||||
|
|
||||||
|
want := alertWebhookURL + ": is not an http or https URL without a user or " +
|
||||||
|
"a fragment, such as https://alerts.example/smallwebwaf"
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSlackAndNtfySettingsAreLoggedWithoutTheirSecrets(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const token = "tk_0123456789abcdefghijklmnopq"
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
alertSlackWebhookURL: "https://hooks.slack.example/services/T0123/B4567/abcdef",
|
||||||
|
alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
|
||||||
|
alertNtfyToken: token,
|
||||||
|
})
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
logged := out.String()
|
||||||
|
for _, want := range []string{
|
||||||
|
`"` + alertSlackWebhookURL + `":"https://hooks.slack.example/********"`,
|
||||||
|
`"` + alertNtfyURL + `":"https://ntfy.example/********"`,
|
||||||
|
`"` + alertNtfyToken + `":"********"`,
|
||||||
|
} {
|
||||||
|
if !strings.Contains(logged, want) {
|
||||||
|
t.Errorf("no %s in the settings logged: %s", want, logged)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, secret := range []string{"T0123", "smallwebwaf-alerts", token} {
|
||||||
|
if strings.Contains(logged, secret) {
|
||||||
|
t.Errorf("%s in the settings logged: %s", secret, logged)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNtfyTokenWithAControlCharacterStopsTheStartWithoutShowingIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// A file saved with Windows line ends keeps the carriage return.
|
||||||
|
for name, env := range map[string]environment{
|
||||||
|
"set": {alertNtfyToken: token + "\r"},
|
||||||
|
"in a file": {alertNtfyToken + "_FILE": writeFile(t, token+"\r\n")},
|
||||||
|
} {
|
||||||
|
_, err := config.FromEnvironment(env.lookupEnv)
|
||||||
|
|
||||||
|
want := alertNtfyToken + ": holds a control character, such as the " +
|
||||||
|
"carriage return of a Windows line end"
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("%s: error %v, want %s", name, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstanceNameWithAControlCharacterStopsTheStartOnlyWithNtfySet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const name = "fsn1app1\r"
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
instanceName: name, alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want := instanceName + `: "fsn1app1\r" holds a control character, such as the ` +
|
||||||
|
`carriage return of a Windows line end, and is sent to ntfy in a header ` +
|
||||||
|
`while ` + alertNtfyURL + ` is set`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{instanceName: name})
|
||||||
|
if cfg.InstanceName != name {
|
||||||
|
t.Errorf("not sending to ntfy, %s is %q", instanceName, cfg.InstanceName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstanceNameNotUTF8StopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// café saved in Latin-1.
|
||||||
|
const latin1 = "caf\xe9"
|
||||||
|
|
||||||
|
for name, env := range map[string]environment{
|
||||||
|
"set": {instanceName: latin1},
|
||||||
|
"in a file": {instanceName + "_FILE": writeFile(t, latin1+"\n")},
|
||||||
|
} {
|
||||||
|
_, err := config.FromEnvironment(env.lookupEnv)
|
||||||
|
|
||||||
|
want := instanceName + `: "caf\xe9" is not valid UTF-8`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("%s: error %v, want %s", name, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -494,6 +835,178 @@ func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLookupSettingsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
lookupTimeout: "500ms", addLookupHeaders: enabled,
|
||||||
|
})
|
||||||
|
wantLookupSettings(t, cfg, config.Config{
|
||||||
|
LookupSource: defaultLookupSource, LookupTimeout: 500 * time.Millisecond,
|
||||||
|
AddLookupHeaders: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
cfg = fromEnvironment(t, environment{lookupSource: off})
|
||||||
|
wantLookupSettings(t, cfg, config.Config{
|
||||||
|
LookupSource: off, LookupTimeout: time.Second,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupDBPathGoesWithTheFileSourceAlone(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const path = "/var/lib/ipinfo/ipinfo_lite.mmdb"
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{lookupSource: fileSource, lookupDBPath: path})
|
||||||
|
if cfg.LookupSource != fileSource || cfg.LookupDBPath != path {
|
||||||
|
t.Errorf("lookups from %q in %q, want file in %q",
|
||||||
|
cfg.LookupSource, cfg.LookupDBPath, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
env environment
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
environment{lookupSource: fileSource},
|
||||||
|
lookupSource + ": is file while " + lookupDBPath +
|
||||||
|
" is unset; it names the file to look clients up in",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
environment{lookupSource: fileSource, lookupDBPath: ""},
|
||||||
|
lookupSource + ": is file while " + lookupDBPath +
|
||||||
|
" is unset; it names the file to look clients up in",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
environment{lookupDBPath: path},
|
||||||
|
lookupDBPath + ": is set while " + lookupSource + " is geojs; only file reads it",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
environment{lookupSource: off, lookupDBPath: path},
|
||||||
|
lookupDBPath + ": is set while " + lookupSource + " is off; only file reads it",
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
_, err := config.FromEnvironment(tc.env.lookupEnv)
|
||||||
|
if err == nil || err.Error() != tc.want {
|
||||||
|
t.Errorf("settings %v: error %v, want %s", tc.env, err, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSettingNeedingLookupsStopsTheStartWhileTheyAreOff(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for name, value := range map[string]string{
|
||||||
|
deniedCountries: "kp",
|
||||||
|
allowedCountries: "de",
|
||||||
|
addLookupHeaders: enabled,
|
||||||
|
asnLimitPercent: "AS64496:50",
|
||||||
|
countryLimitPercent: "cn:25",
|
||||||
|
asnBytesPercent: "AS64496:50",
|
||||||
|
countryBytesPercent: "cn:25",
|
||||||
|
unknownLimitPercent: "99",
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
lookupSource: off, name: value,
|
||||||
|
}.lookupEnv)
|
||||||
|
if err == nil || !strings.HasPrefix(err.Error(), name+": ") ||
|
||||||
|
!strings.Contains(err.Error(), lookupSource+" is off") {
|
||||||
|
t.Errorf("error %v, want one naming %s and %s", err, name, lookupSource)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set empty, the lists need nothing looked up, and nor does
|
||||||
|
// SWWAF_UNKNOWN_LIMIT_PERCENT at 100, which lowers no limit.
|
||||||
|
fromEnvironment(t, environment{
|
||||||
|
lookupSource: off, deniedCountries: "", allowedCountries: "",
|
||||||
|
asnLimitPercent: "", countryLimitPercent: "", asnBytesPercent: "",
|
||||||
|
countryBytesPercent: "", unknownLimitPercent: "100",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBiasedThresholdsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{})
|
||||||
|
if len(cfg.ASNLimitPercent) != 0 || len(cfg.CountryLimitPercent) != 0 ||
|
||||||
|
len(cfg.ASNBytesPercent) != 0 || len(cfg.CountryBytesPercent) != 0 ||
|
||||||
|
cfg.UnknownLimitPercent != 100 {
|
||||||
|
t.Errorf("biased thresholds %v, %v, %v, %v and %d by default, "+
|
||||||
|
"want four empty lists and 100", cfg.ASNLimitPercent, cfg.CountryLimitPercent,
|
||||||
|
cfg.ASNBytesPercent, cfg.CountryBytesPercent, cfg.UnknownLimitPercent)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AS numbers and countries in either case, an AS number with leading
|
||||||
|
// zeros, 0 and 100.
|
||||||
|
cfg = fromEnvironment(t, environment{
|
||||||
|
asnLimitPercent: "AS14061:50, as16276:0,AS045102:100",
|
||||||
|
countryLimitPercent: "cn:25,RU:50",
|
||||||
|
asnBytesPercent: "as16276:75",
|
||||||
|
countryBytesPercent: "ru:10",
|
||||||
|
unknownLimitPercent: "0",
|
||||||
|
})
|
||||||
|
|
||||||
|
for name, tc := range map[string]struct{ got, want map[string]int64 }{
|
||||||
|
asnLimitPercent: {
|
||||||
|
cfg.ASNLimitPercent,
|
||||||
|
map[string]int64{"AS14061": 50, "AS16276": 0, "AS45102": 100},
|
||||||
|
},
|
||||||
|
countryLimitPercent: {cfg.CountryLimitPercent, map[string]int64{"CN": 25, "RU": 50}},
|
||||||
|
asnBytesPercent: {cfg.ASNBytesPercent, map[string]int64{"AS16276": 75}},
|
||||||
|
countryBytesPercent: {cfg.CountryBytesPercent, map[string]int64{"RU": 10}},
|
||||||
|
} {
|
||||||
|
if !maps.Equal(tc.got, tc.want) {
|
||||||
|
t.Errorf("%s gave %v, want %v", name, tc.got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.UnknownLimitPercent != 0 {
|
||||||
|
t.Errorf("%s gave %d, want 0", unknownLimitPercent, cfg.UnknownLimitPercent)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInvalidBiasedThresholdStopsTheStartSayingWhatIsWrong(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
notASN = " is not an AS number such as AS64496"
|
||||||
|
notItem = " is not a code, : and a percentage, such as AS64496:50 or cn:25"
|
||||||
|
notPercent = " is not a percentage, a whole number from 0 to 100"
|
||||||
|
)
|
||||||
|
|
||||||
|
for _, tc := range []struct{ name, value, want string }{
|
||||||
|
{asnLimitPercent, "14061:50", `"14061"` + notASN},
|
||||||
|
{asnLimitPercent, "AS4294967296:50", `"AS4294967296"` + notASN},
|
||||||
|
{asnLimitPercent, "AS14061", `"AS14061"` + notItem},
|
||||||
|
{asnLimitPercent, "AS14061:101", `"101"` + notPercent},
|
||||||
|
{asnLimitPercent, "AS14061:50,as14061:25", `"as14061" is listed twice`},
|
||||||
|
{
|
||||||
|
countryLimitPercent, "nk:25",
|
||||||
|
`"nk" is not a two-letter country code such as de or kp`,
|
||||||
|
},
|
||||||
|
{countryLimitPercent, "cn:25,CN:50", `"CN" is listed twice`},
|
||||||
|
{asnBytesPercent, "AS14061:-1", `"-1"` + notPercent},
|
||||||
|
{countryBytesPercent, "cn:50%", `"50%"` + notPercent},
|
||||||
|
{unknownLimitPercent, "101", `"101"` + notPercent},
|
||||||
|
{unknownLimitPercent, off, `"off"` + notPercent},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
|
||||||
|
|
||||||
|
want := tc.name + ": " + tc.want
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestSizesAndOff(t *testing.T) {
|
func TestSizesAndOff(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -615,6 +1128,12 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{rateLimitPerHour, "1.5"},
|
{rateLimitPerHour, "1.5"},
|
||||||
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
|
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
|
||||||
{rateLimitExemptPaths, "/assets/,,/static/"},
|
{rateLimitExemptPaths, "/assets/,,/static/"},
|
||||||
|
{bytesLimitPerMinute, "10GB"}, {bytesLimitPerHour, "0"},
|
||||||
|
{bytesLimitPerDay, "-1G"},
|
||||||
|
{bytesCount, "all"}, {bytesCount, "Both"}, {bytesCount, ""},
|
||||||
|
{lookupSource, "ipinfo"}, {lookupSource, "GeoJS"}, {lookupSource, ""},
|
||||||
|
{lookupTimeout, off}, {lookupTimeout, "0s"}, {lookupTimeout, "1"},
|
||||||
|
{addLookupHeaders, "yes"},
|
||||||
{deniedCountries, "nk"},
|
{deniedCountries, "nk"},
|
||||||
{deniedCountries, "kp,,ir"},
|
{deniedCountries, "kp,,ir"},
|
||||||
{deniedCountries, "prk"},
|
{deniedCountries, "prk"},
|
||||||
@@ -627,6 +1146,14 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{allowedCountries, "uk"},
|
{allowedCountries, "uk"},
|
||||||
{allowedCountries, "zz"},
|
{allowedCountries, "zz"},
|
||||||
{allowedCountries, "de,germany"},
|
{allowedCountries, "de,germany"},
|
||||||
|
{asnLimitPercent, "AS14061:50,,AS16276:50"}, {asnLimitPercent, "ASX:50"},
|
||||||
|
{asnLimitPercent, "AS14061:"}, {asnLimitPercent, "AS14061 :50"},
|
||||||
|
{asnLimitPercent, "AS14061:1.5"}, {asnLimitPercent, "AS-1:50"},
|
||||||
|
{countryLimitPercent, "cn"}, {countryLimitPercent, "cn:"},
|
||||||
|
{countryLimitPercent, "cn:25:50"}, {countryLimitPercent, "china:25"},
|
||||||
|
{asnBytesPercent, "AS14061:101"}, {countryBytesPercent, "su:50"},
|
||||||
|
{unknownLimitPercent, ""}, {unknownLimitPercent, "-1"},
|
||||||
|
{unknownLimitPercent, "50%"},
|
||||||
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
|
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
|
||||||
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
|
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
|
||||||
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
|
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
|
||||||
@@ -866,20 +1393,6 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"})
|
cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"})
|
||||||
|
|
||||||
var out bytes.Buffer
|
|
||||||
|
|
||||||
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
|
||||||
|
|
||||||
var line struct {
|
|
||||||
Settings map[string]string `json:"settings"`
|
|
||||||
}
|
|
||||||
|
|
||||||
err := json.Unmarshal(out.Bytes(), &line)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
|
||||||
}
|
|
||||||
|
|
||||||
hostname, _ := os.Hostname()
|
hostname, _ := os.Hostname()
|
||||||
|
|
||||||
want := map[string]string{
|
want := map[string]string{
|
||||||
@@ -902,8 +1415,21 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
rateLimitPerHour: "10000",
|
rateLimitPerHour: "10000",
|
||||||
rateLimitPerDay: "50000",
|
rateLimitPerDay: "50000",
|
||||||
rateLimitExemptPaths: "",
|
rateLimitExemptPaths: "",
|
||||||
|
bytesLimitPerMinute: "10G",
|
||||||
|
bytesLimitPerHour: "20G",
|
||||||
|
bytesLimitPerDay: "50G",
|
||||||
|
bytesCount: "both",
|
||||||
|
lookupSource: defaultLookupSource,
|
||||||
|
lookupDBPath: "",
|
||||||
|
lookupTimeout: "1s",
|
||||||
|
addLookupHeaders: "false",
|
||||||
deniedCountries: "",
|
deniedCountries: "",
|
||||||
allowedCountries: "",
|
allowedCountries: "",
|
||||||
|
asnLimitPercent: "",
|
||||||
|
countryLimitPercent: "",
|
||||||
|
asnBytesPercent: "",
|
||||||
|
countryBytesPercent: "",
|
||||||
|
unknownLimitPercent: "100",
|
||||||
banResponse: "403",
|
banResponse: "403",
|
||||||
limitBanDuration: "1h",
|
limitBanDuration: "1h",
|
||||||
limitBanRepeatWindow: "24h",
|
limitBanRepeatWindow: "24h",
|
||||||
@@ -926,12 +1452,40 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
logRemoteBuffer: "10000",
|
logRemoteBuffer: "10000",
|
||||||
logRemoteFacility: "local0",
|
logRemoteFacility: "local0",
|
||||||
logRemoteAppName: hostname,
|
logRemoteAppName: hostname,
|
||||||
|
alertWebhookURL: "",
|
||||||
|
alertWebhookHeaders: "",
|
||||||
|
alertSlackWebhookURL: "",
|
||||||
|
alertNtfyURL: "",
|
||||||
|
alertNtfyToken: "",
|
||||||
|
alertEvents: defaultAlertEvents,
|
||||||
|
alertCooldown: defaultAlertCooldown,
|
||||||
|
alertMaxPerHour: "60",
|
||||||
}
|
}
|
||||||
if !maps.Equal(line.Settings, want) {
|
if got := loggedSettings(t, cfg); !maps.Equal(got, want) {
|
||||||
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
|
t.Errorf("logged settings\n%v\nwant\n%v", got, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// loggedSettings returns the settings as cfg logs them, each by its name.
|
||||||
|
func loggedSettings(t *testing.T, cfg *config.Config) map[string]string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
var line struct {
|
||||||
|
Settings map[string]string `json:"settings"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err := json.Unmarshal(out.Bytes(), &line)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return line.Settings
|
||||||
|
}
|
||||||
|
|
||||||
// wantSettings checks the settings that are plain values.
|
// wantSettings checks the settings that are plain values.
|
||||||
func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
@@ -955,6 +1509,33 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
|||||||
wantBanSettings(t, got, want)
|
wantBanSettings(t, got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// wantByteLimitSettings checks the settings for the byte limits.
|
||||||
|
func wantByteLimitSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if got.BytesLimitPerMinute != want.BytesLimitPerMinute ||
|
||||||
|
got.BytesLimitPerHour != want.BytesLimitPerHour ||
|
||||||
|
got.BytesLimitPerDay != want.BytesLimitPerDay ||
|
||||||
|
got.BytesCount != want.BytesCount {
|
||||||
|
t.Errorf("byte limits %d, %d and %d counting %s, want %d, %d and %d counting %s",
|
||||||
|
got.BytesLimitPerMinute, got.BytesLimitPerHour, got.BytesLimitPerDay,
|
||||||
|
got.BytesCount, want.BytesLimitPerMinute, want.BytesLimitPerHour,
|
||||||
|
want.BytesLimitPerDay, want.BytesCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantLookupSettings checks the settings for lookups.
|
||||||
|
func wantLookupSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if got.LookupSource != want.LookupSource || got.LookupTimeout != want.LookupTimeout ||
|
||||||
|
got.AddLookupHeaders != want.AddLookupHeaders {
|
||||||
|
t.Errorf("lookups from %q, waited for %s, headers %t; want %q, %s, %t",
|
||||||
|
got.LookupSource, got.LookupTimeout, got.AddLookupHeaders,
|
||||||
|
want.LookupSource, want.LookupTimeout, want.AddLookupHeaders)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// wantBanSettings checks the settings for bans, the state files, the
|
// wantBanSettings checks the settings for bans, the state files, the
|
||||||
// metrics and the rule files.
|
// metrics and the rule files.
|
||||||
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
|
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
package lookup
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fsnotify/fsnotify"
|
||||||
|
"github.com/oschwald/maxminddb-golang/v2"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
)
|
||||||
|
|
||||||
|
// quietTime is how long the lookup database must go without a change
|
||||||
|
// before it is read again, so that a file still being copied in is read
|
||||||
|
// only once whole.
|
||||||
|
const quietTime = 2 * time.Second
|
||||||
|
|
||||||
|
// FileParams are what OpenFile needs.
|
||||||
|
type FileParams struct {
|
||||||
|
// Path is the lookup database, the IPinfo Lite file in its .mmdb form
|
||||||
|
// (SWWAF_LOOKUP_DB_PATH).
|
||||||
|
Path string
|
||||||
|
// Now tells the time, normally time.Now.
|
||||||
|
Now func() time.Time
|
||||||
|
// ProcessLog receives each reading of the file, and why a replacement
|
||||||
|
// of it cannot be read.
|
||||||
|
ProcessLog *slog.Logger
|
||||||
|
// Alerts receive a file_error alert for each replacement that cannot
|
||||||
|
// be read.
|
||||||
|
Alerts *alerts.Queue
|
||||||
|
}
|
||||||
|
|
||||||
|
// File looks up clients' AS numbers and countries in the lookup database,
|
||||||
|
// held in memory, and reads it again when it is replaced. It is safe for
|
||||||
|
// concurrent use.
|
||||||
|
type File struct {
|
||||||
|
params FileParams
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
// reader is the database in use, and lastRead when it was read.
|
||||||
|
// readFailures are the replacements that could not be read.
|
||||||
|
reader *maxminddb.Reader
|
||||||
|
lastRead time.Time
|
||||||
|
readFailures int
|
||||||
|
}
|
||||||
|
|
||||||
|
// record is what the lookup database holds about a network, of the fields
|
||||||
|
// smallwebwaf reads.
|
||||||
|
type record struct {
|
||||||
|
ASN string `maxminddb:"asn"`
|
||||||
|
ASName string `maxminddb:"as_name"`
|
||||||
|
CountryCode string `maxminddb:"country_code"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenFile reads the lookup database. A file that cannot be read, or that
|
||||||
|
// is not a .mmdb file, is an error.
|
||||||
|
func OpenFile(params FileParams) (*File, error) {
|
||||||
|
reader, err := read(params.Path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
f := &File{params: params}
|
||||||
|
f.use(reader)
|
||||||
|
|
||||||
|
return f, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LookUp returns what the lookup database says about client: its AS
|
||||||
|
// number, such as AS64496, the AS's name, and its country, such as DE,
|
||||||
|
// each "" when the database does not give it, as for an address missing
|
||||||
|
// from it. The database is asked about the client's first address, as
|
||||||
|
// GeoJS is.
|
||||||
|
func (f *File) LookUp(client netip.Prefix) Answer {
|
||||||
|
f.mu.Lock()
|
||||||
|
reader := f.reader
|
||||||
|
f.mu.Unlock()
|
||||||
|
|
||||||
|
var found record
|
||||||
|
|
||||||
|
// A record that cannot be decoded places the client nowhere, as a
|
||||||
|
// missing one does.
|
||||||
|
err := reader.Lookup(client.Addr()).Decode(&found)
|
||||||
|
if err != nil {
|
||||||
|
found = record{}
|
||||||
|
}
|
||||||
|
|
||||||
|
return Answer{
|
||||||
|
Client: client,
|
||||||
|
ASN: found.ASN,
|
||||||
|
ASName: found.ASName,
|
||||||
|
Country: found.CountryCode,
|
||||||
|
Answered: f.params.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// LastRead returns when the lookup database in use was read.
|
||||||
|
func (f *File) LastRead() time.Time {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
return f.lastRead
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadFailures returns how many replacements of the lookup database could
|
||||||
|
// not be read.
|
||||||
|
func (f *File) ReadFailures() int {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
return f.readFailures
|
||||||
|
}
|
||||||
|
|
||||||
|
// Watch watches the directory of the lookup database until ctx is done,
|
||||||
|
// and reads the file again once it has gone without a change for
|
||||||
|
// quietTime, after it is replaced, written or removed, and after Watch
|
||||||
|
// starts watching. If the directory cannot be watched, that is logged, and
|
||||||
|
// the database read at start stays in use.
|
||||||
|
func (f *File) Watch(ctx context.Context) {
|
||||||
|
watcher, err := fsnotify.NewWatcher()
|
||||||
|
if err == nil {
|
||||||
|
defer func() {
|
||||||
|
_ = watcher.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
err = watcher.Add(filepath.Dir(f.params.Path))
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error("cannot watch the lookup database for replacements",
|
||||||
|
"error", err.Error())
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.ProcessLog.Info("watching the lookup database for replacements",
|
||||||
|
"file", f.params.Path)
|
||||||
|
|
||||||
|
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readAfterChanges reads the lookup database again once quietTime has
|
||||||
|
// passed without a change to it from events, until ctx is done, and logs
|
||||||
|
// the errors from errs. A change to another file in its directory does not
|
||||||
|
// count. The wait starts at once, as if for a change, so that a file
|
||||||
|
// replaced after OpenFile read it, and before its directory was watched,
|
||||||
|
// is read too.
|
||||||
|
func (f *File) readAfterChanges(
|
||||||
|
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
|
||||||
|
) {
|
||||||
|
path := filepath.Clean(f.params.Path)
|
||||||
|
|
||||||
|
quiet := time.NewTimer(quietTime)
|
||||||
|
defer quiet.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case event := <-events:
|
||||||
|
if filepath.Clean(event.Name) == path {
|
||||||
|
quiet.Reset(quietTime)
|
||||||
|
}
|
||||||
|
case <-quiet.C:
|
||||||
|
f.readAgain()
|
||||||
|
case err := <-errs:
|
||||||
|
f.params.ProcessLog.Warn("watching the lookup database failed",
|
||||||
|
"error", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// readAgain reads the lookup database again, in place of the one in use,
|
||||||
|
// or, if it cannot be read, counts that, raises a file_error alert for it
|
||||||
|
// and logs it, and the one in use stays in use.
|
||||||
|
func (f *File) readAgain() {
|
||||||
|
reader, err := read(f.params.Path)
|
||||||
|
if err != nil {
|
||||||
|
const kept = "the lookup database cannot be read, " +
|
||||||
|
"and the one read before stays in use"
|
||||||
|
|
||||||
|
f.mu.Lock()
|
||||||
|
f.readFailures++
|
||||||
|
f.mu.Unlock()
|
||||||
|
|
||||||
|
// Raised before it is logged, so that the alert is there once the
|
||||||
|
// log line is.
|
||||||
|
f.params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventFileError,
|
||||||
|
Reason: kept,
|
||||||
|
Detail: map[string]any{"file": f.params.Path, "error": err.Error()},
|
||||||
|
})
|
||||||
|
f.params.ProcessLog.Error(kept, "error", err.Error())
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f.use(reader)
|
||||||
|
}
|
||||||
|
|
||||||
|
// use puts reader in use, in place of the database read before, and logs
|
||||||
|
// that the file was read.
|
||||||
|
func (f *File) use(reader *maxminddb.Reader) {
|
||||||
|
f.mu.Lock()
|
||||||
|
f.reader = reader
|
||||||
|
f.lastRead = f.params.Now()
|
||||||
|
f.mu.Unlock()
|
||||||
|
|
||||||
|
f.params.ProcessLog.Info("read the lookup database", "file", f.params.Path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// read reads the lookup database at path. The whole file is read into
|
||||||
|
// memory, rather than mapped into it as the reader can, so that a file
|
||||||
|
// overwritten in place cannot change, or end, under a lookup.
|
||||||
|
func read(path string) (*maxminddb.Reader, error) {
|
||||||
|
data, err := os.ReadFile(path) //nolint:gosec // the file the admin names
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH cannot be read: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
reader, err := maxminddb.OpenBytes(data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH %s is not a .mmdb file: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return reader, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,379 @@
|
|||||||
|
package lookup
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fsnotify/fsnotify"
|
||||||
|
"github.com/maxmind/mmdbwriter/mmdbtype"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testNetblock is the netblock the tests' lookup databases place, and
|
||||||
|
// testClient a client in it.
|
||||||
|
const (
|
||||||
|
testNetblock = "203.0.113.0/24"
|
||||||
|
testClient = "203.0.113.9/32"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFilePlacesClientsAndCountsAnAddressMissingFromItAsUnknown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
germany := lookuptest.Network{ASN: "AS64496", ASName: "Example Net", Country: "DE"}
|
||||||
|
northKorea := lookuptest.Network{ASN: "AS64511", ASName: "Other Net", Country: "KP"}
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||||
|
testNetblock: germany,
|
||||||
|
"2001:db8::/32": northKorea,
|
||||||
|
})
|
||||||
|
|
||||||
|
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||||
|
|
||||||
|
f, err := OpenFile(FileParams{
|
||||||
|
Path: path,
|
||||||
|
Now: func() time.Time { return now },
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Alerts: newQueue(),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for client, want := range map[string]lookuptest.Network{
|
||||||
|
testClient: germany,
|
||||||
|
// An IPv6 client is its /64.
|
||||||
|
"2001:db8:1:2::/64": northKorea,
|
||||||
|
"198.51.100.7/32": {},
|
||||||
|
} {
|
||||||
|
prefix := netip.MustParsePrefix(client)
|
||||||
|
|
||||||
|
got := f.LookUp(prefix)
|
||||||
|
if got != (Answer{
|
||||||
|
Client: prefix, ASN: want.ASN, ASName: want.ASName, Country: want.Country,
|
||||||
|
Answered: now,
|
||||||
|
}) {
|
||||||
|
t.Errorf("%s has the answer %+v, want %+v, answered %s", client, got, want, now)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordThatCannotBeReadPlacesTheClientNowhere(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// The AS number is a number, where a string belongs. The writer writes
|
||||||
|
// a record's fields in the order of their names, so as_name is read
|
||||||
|
// before the AS number fails.
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
lookuptest.WriteRecords(t, path, map[string]mmdbtype.Map{
|
||||||
|
testNetblock: {
|
||||||
|
"asn": mmdbtype.Uint32(64496),
|
||||||
|
"as_name": mmdbtype.String("Example Net"),
|
||||||
|
"country_code": mmdbtype.String("DE"),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
f := openFile(t, path, newQueue())
|
||||||
|
|
||||||
|
answer := f.LookUp(netip.MustParsePrefix(testClient))
|
||||||
|
if answer.ASN != "" || answer.ASName != "" || answer.Country != "" {
|
||||||
|
t.Errorf("%s is placed %+v, want nowhere", testClient, answer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFileThatCannotBeReadIsAnError(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
missing := filepath.Join(dir, "missing.mmdb")
|
||||||
|
notDatabase := filepath.Join(dir, "not.mmdb")
|
||||||
|
writeFile(t, notDatabase, "not a lookup database\n")
|
||||||
|
|
||||||
|
for path, want := range map[string]string{
|
||||||
|
missing: "SWWAF_LOOKUP_DB_PATH cannot be read: open " + missing +
|
||||||
|
": no such file or directory",
|
||||||
|
notDatabase: "SWWAF_LOOKUP_DB_PATH " + notDatabase +
|
||||||
|
" is not a .mmdb file: error opening database: invalid MaxMind DB file",
|
||||||
|
} {
|
||||||
|
_, err := OpenFile(FileParams{
|
||||||
|
Path: path,
|
||||||
|
Now: time.Now,
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Alerts: newQueue(),
|
||||||
|
})
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("opening %s failed with %v, want %s", path, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The tests below run readAfterChanges in a synctest bubble, where time is
|
||||||
|
// a clock of the test's own: time.Sleep moves it on at once, and
|
||||||
|
// synctest.Wait returns once readAfterChanges waits again, so that every
|
||||||
|
// reading due by then is done. The test sends the changes itself, as the
|
||||||
|
// watch of a directory cannot run in a bubble.
|
||||||
|
|
||||||
|
func TestReplacementCopiedOverTheFileInTwoPartsIsReadOnlyWhole(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
||||||
|
writeDatabase(t, path, "DE")
|
||||||
|
|
||||||
|
queue := newQueue()
|
||||||
|
f := openFile(t, path, queue)
|
||||||
|
changes := watch(t, f)
|
||||||
|
|
||||||
|
other := filepath.Join(dir, "replacement.mmdb")
|
||||||
|
writeDatabase(t, other, "KP")
|
||||||
|
|
||||||
|
replacement, err := os.ReadFile(other) //nolint:gosec // a file the test wrote
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read %s: %v", other, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The file in use is overwritten in place, and keeps giving what
|
||||||
|
// it gave. Its first part alone is not a .mmdb file.
|
||||||
|
file, err := os.Create(path) //nolint:gosec // a file the test wrote
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = file.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
half := len(replacement) / 2
|
||||||
|
write(t, file, replacement[:half])
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
time.Sleep(quietTime - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "DE")
|
||||||
|
|
||||||
|
// The second part starts the wait again.
|
||||||
|
write(t, file, replacement[half:])
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
time.Sleep(quietTime - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "DE")
|
||||||
|
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "KP")
|
||||||
|
|
||||||
|
if !f.LastRead().Equal(time.Now()) || f.ReadFailures() != 0 {
|
||||||
|
t.Errorf("read at %s, with %d failures; want read now, with none",
|
||||||
|
f.LastRead(), f.ReadFailures())
|
||||||
|
}
|
||||||
|
|
||||||
|
wantAlerts(t, queue)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReplacementThatCannotBeReadLeavesTheFileInUseWithOneAlert(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
writeDatabase(t, path, "DE")
|
||||||
|
|
||||||
|
queue := newQueue()
|
||||||
|
f := openFile(t, path, queue)
|
||||||
|
read := f.LastRead()
|
||||||
|
changes := watch(t, f)
|
||||||
|
|
||||||
|
writeFile(t, path, "not a lookup database\n")
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
// Long after, the replacement has been read once.
|
||||||
|
time.Sleep(time.Hour)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "DE")
|
||||||
|
|
||||||
|
if !f.LastRead().Equal(read) || f.ReadFailures() != 1 {
|
||||||
|
t.Errorf("read at %s, with %d failures; want read at %s, with one",
|
||||||
|
f.LastRead(), f.ReadFailures(), read)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantAlerts(t, queue, alerts.Alert{
|
||||||
|
Time: read.Add(quietTime),
|
||||||
|
Event: alerts.EventFileError,
|
||||||
|
Reason: "the lookup database cannot be read, and the one read before stays in use",
|
||||||
|
Detail: map[string]any{
|
||||||
|
"file": path,
|
||||||
|
"error": "SWWAF_LOOKUP_DB_PATH " + path + " is not a .mmdb file: " +
|
||||||
|
"error opening database: invalid MaxMind DB file",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestChangeOfAnotherFileInTheDirectoryIsNoReplacement(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
||||||
|
writeDatabase(t, path, "DE")
|
||||||
|
|
||||||
|
f := openFile(t, path, newQueue())
|
||||||
|
changes := watch(t, f)
|
||||||
|
|
||||||
|
// The wait that starts with the watch ends with a reading.
|
||||||
|
time.Sleep(quietTime)
|
||||||
|
synctest.Wait()
|
||||||
|
writeDatabase(t, path, "KP")
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: filepath.Join(dir, "other.mmdb"), Op: fsnotify.Create}
|
||||||
|
|
||||||
|
time.Sleep(quietTime)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "DE")
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
time.Sleep(quietTime)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "KP")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReplacementSavedBeforeTheWatchStartsIsRead(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
writeDatabase(t, path, "DE")
|
||||||
|
|
||||||
|
f := openFile(t, path, newQueue())
|
||||||
|
|
||||||
|
// Saved after OpenFile read the file, and before its directory was
|
||||||
|
// watched, so that no change is seen for it.
|
||||||
|
writeDatabase(t, path, "KP")
|
||||||
|
watch(t, f)
|
||||||
|
time.Sleep(quietTime)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCountry(t, f, "KP")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// newQueue returns a queue of alerts for a webhook that is never sent
|
||||||
|
// them, so that they wait in it for the test to look at.
|
||||||
|
func newQueue() *alerts.Queue {
|
||||||
|
return alerts.New(alerts.Params{
|
||||||
|
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||||
|
Events: alerts.Events(),
|
||||||
|
Cooldown: 15 * time.Minute,
|
||||||
|
Now: time.Now,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeDatabase writes a lookup database at path that places testNetblock
|
||||||
|
// in country, and no other address.
|
||||||
|
func writeDatabase(t *testing.T, path, country string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||||
|
testNetblock: {ASN: "AS64496", ASName: "Example Net", Country: country},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// openFile opens the lookup database at path, which raises its alerts to
|
||||||
|
// queue.
|
||||||
|
func openFile(t *testing.T, path string, queue *alerts.Queue) *File {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
f, err := OpenFile(FileParams{
|
||||||
|
Path: path,
|
||||||
|
Now: time.Now,
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Alerts: queue,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
// watch runs f's readAfterChanges until the test ends, and returns the
|
||||||
|
// channel that sends it changes.
|
||||||
|
func watch(t *testing.T, f *File) chan<- fsnotify.Event {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
changes := make(chan fsnotify.Event)
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
f.readAfterChanges(ctx, changes, nil)
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
stop()
|
||||||
|
<-stopped
|
||||||
|
})
|
||||||
|
|
||||||
|
return changes
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantCountry checks the country f gives testClient.
|
||||||
|
func wantCountry(t *testing.T, f *File, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := f.LookUp(netip.MustParsePrefix(testClient)).Country
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("%s is in %q, want %q", testClient, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantAlerts checks the alerts waiting in queue, and that it held none
|
||||||
|
// back.
|
||||||
|
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != len(want) || (len(want) > 0 && !reflect.DeepEqual(waiting, want)) {
|
||||||
|
t.Errorf("alerts waiting %+v, want %+v", waiting, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if queue.Suppressed() != 0 {
|
||||||
|
t.Errorf("%d alerts held back, want none", queue.Suppressed())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeFile writes content to the file at path.
|
||||||
|
func writeFile(t *testing.T, path, content string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// write writes data to the end of file.
|
||||||
|
func write(t *testing.T, file *os.File, data []byte) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
_, err := file.Write(data)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+140
-65
@@ -1,7 +1,8 @@
|
|||||||
// Package lookup looks up each client's country through the GeoJS web
|
// Package lookup looks up each client's AS number and country, through
|
||||||
// service, and keeps the answers in memory, for at most 100,000 clients
|
// the GeoJS web service or in the lookup database, the IPinfo Lite file
|
||||||
// and for 7 days each. The answers are written to lookups.json and read
|
// SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at
|
||||||
// from it by the state package.
|
// most 100,000 clients and for 7 days each, and are written to
|
||||||
|
// lookups.json and read from it by the state package.
|
||||||
package lookup
|
package lookup
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -14,17 +15,20 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
// URL is GeoJS's country endpoint. Asked about several addresses at once,
|
// URL is GeoJS's endpoint for an address's place and network. Asked about
|
||||||
// comma separated in its ip parameter, it answers with a list.
|
// several addresses at once, comma separated in its ip parameter, it
|
||||||
const URL = "https://get.geojs.io/v1/ip/country.json"
|
// answers with a list.
|
||||||
|
const URL = "https://get.geojs.io/v1/ip/geo.json"
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// keepFor is how long an answer is used instead of asking GeoJS again.
|
// keepFor is how long an answer is used instead of asking GeoJS again.
|
||||||
@@ -39,9 +43,8 @@ const (
|
|||||||
maxWaiting = 10000
|
maxWaiting = 10000
|
||||||
// maxPerRequest is how many addresses one request to GeoJS asks about.
|
// maxPerRequest is how many addresses one request to GeoJS asks about.
|
||||||
maxPerRequest = 200
|
maxPerRequest = 200
|
||||||
// timeout is how long a new client waits for its answer, and how long
|
// unknownASN is the AS number GeoJS gives when it knows none.
|
||||||
// a request to GeoJS may take before it is abandoned.
|
unknownASN = 64512
|
||||||
timeout = time.Second
|
|
||||||
// After a failure GeoJS is not asked again for a second, and for
|
// After a failure GeoJS is not asked again for a second, and for
|
||||||
// retryDelayFactor times as long after each further failure in a row,
|
// retryDelayFactor times as long after each further failure in a row,
|
||||||
// up to five minutes.
|
// up to five minutes.
|
||||||
@@ -61,6 +64,16 @@ var (
|
|||||||
type Params struct {
|
type Params struct {
|
||||||
// URL is where GeoJS is asked, normally URL.
|
// URL is where GeoJS is asked, normally URL.
|
||||||
URL string
|
URL string
|
||||||
|
// Timeout is how long a request waits for its client's first answer,
|
||||||
|
// and how long a request to GeoJS may take before it is abandoned
|
||||||
|
// (SWWAF_LOOKUP_TIMEOUT).
|
||||||
|
Timeout time.Duration
|
||||||
|
// Wait is true when a setting needs each request's answer before the
|
||||||
|
// request goes on. Otherwise no request waits for one.
|
||||||
|
Wait bool
|
||||||
|
// Answered, unless nil, is given each answer GeoJS gives, once it is
|
||||||
|
// kept.
|
||||||
|
Answered func(Answer)
|
||||||
// Now tells the time, normally time.Now.
|
// Now tells the time, normally time.Now.
|
||||||
Now func() time.Time
|
Now func() time.Time
|
||||||
// ProcessLog receives GeoJS's failures.
|
// ProcessLog receives GeoJS's failures.
|
||||||
@@ -68,16 +81,22 @@ type Params struct {
|
|||||||
// Metrics count the requests to GeoJS, those that failed, and the
|
// Metrics count the requests to GeoJS, those that failed, and the
|
||||||
// clients that go without an answer.
|
// clients that go without an answer.
|
||||||
Metrics *metrics.Metrics
|
Metrics *metrics.Metrics
|
||||||
|
// Alerts receive a source_failure alert each time GeoJS fails.
|
||||||
|
Alerts *alerts.Queue
|
||||||
}
|
}
|
||||||
|
|
||||||
// GeoJS looks up clients' countries through GeoJS. At most one request
|
// GeoJS looks up clients' AS numbers and countries through GeoJS. At most
|
||||||
// to GeoJS is under way at a time, and it asks about every client waiting,
|
// one request to GeoJS is under way at a time, and it asks about every
|
||||||
// up to maxPerRequest. It is safe for concurrent use.
|
// client waiting, up to maxPerRequest. It is safe for concurrent use.
|
||||||
type GeoJS struct {
|
type GeoJS struct {
|
||||||
url string
|
url string
|
||||||
|
timeout time.Duration
|
||||||
|
wait bool
|
||||||
|
answered func(Answer)
|
||||||
now func() time.Time
|
now func() time.Time
|
||||||
processLog *slog.Logger
|
processLog *slog.Logger
|
||||||
metrics *metrics.Metrics
|
metrics *metrics.Metrics
|
||||||
|
alerts *alerts.Queue
|
||||||
// httpClient follows no redirect, so that visitors' addresses go to
|
// httpClient follows no redirect, so that visitors' addresses go to
|
||||||
// GeoJS alone: a redirect is a failure.
|
// GeoJS alone: a redirect is a failure.
|
||||||
httpClient *http.Client
|
httpClient *http.Client
|
||||||
@@ -95,11 +114,18 @@ type GeoJS struct {
|
|||||||
retryAt time.Time
|
retryAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
// Answer is what GeoJS said about a client, as lookups.json holds it: its
|
// Answer is what GeoJS or the lookup database said about a client: its AS
|
||||||
// country, "" when GeoJS cannot place it, when GeoJS said so, and when
|
// number, such as AS64496, and the AS's name, both "" when the source knows
|
||||||
// the answer was last used.
|
// no AS number for it; its country, "" when the source cannot place it;
|
||||||
|
// when the source said so; and, for GeoJS's answers, which lookups.json
|
||||||
|
// holds, when the answer was last used. The zero Answer is that of a
|
||||||
|
// client with no answer.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
type Answer struct {
|
type Answer struct {
|
||||||
Client netip.Prefix `json:"client"`
|
Client netip.Prefix `json:"client"`
|
||||||
|
ASN string `json:"asn"`
|
||||||
|
ASName string `json:"as_name"`
|
||||||
Country string `json:"country"`
|
Country string `json:"country"`
|
||||||
Answered time.Time `json:"answered"`
|
Answered time.Time `json:"answered"`
|
||||||
Used time.Time `json:"used"`
|
Used time.Time `json:"used"`
|
||||||
@@ -124,9 +150,13 @@ func New(params Params) *GeoJS {
|
|||||||
|
|
||||||
return &GeoJS{
|
return &GeoJS{
|
||||||
url: params.URL,
|
url: params.URL,
|
||||||
|
timeout: params.Timeout,
|
||||||
|
wait: params.Wait,
|
||||||
|
answered: params.Answered,
|
||||||
now: params.Now,
|
now: params.Now,
|
||||||
processLog: params.ProcessLog,
|
processLog: params.ProcessLog,
|
||||||
metrics: params.Metrics,
|
metrics: params.Metrics,
|
||||||
|
alerts: params.Alerts,
|
||||||
httpClient: &http.Client{
|
httpClient: &http.Client{
|
||||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||||
return http.ErrUseLastResponse
|
return http.ErrUseLastResponse
|
||||||
@@ -137,23 +167,23 @@ func New(params Params) *GeoJS {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Country returns the country GeoJS places client in, as a two-letter
|
// LookUp returns the answer GeoJS gave about client, with its country as
|
||||||
// code in capitals, or "" when the country cannot be found: GeoJS cannot
|
// a two-letter code in capitals, or the zero Answer when there is none
|
||||||
// place the client, or has not answered in time. An answer is kept for 7
|
// yet. An answer is kept for 7 days. Without one, the client is asked
|
||||||
// days. Without one, a client waits up to timeout for it, unless it has
|
// about in the background, and, while Wait is set, the request waits up
|
||||||
// gone without one before; until GeoJS answers, the client is asked about
|
// to Timeout for the answer, unless the client has gone without one
|
||||||
// again in the background. ctx is the context of the client's request,
|
// before. ctx is the context of the client's request, and ends the wait
|
||||||
// and ends the wait when it ends.
|
// when it ends.
|
||||||
//
|
//
|
||||||
// GeoJS is asked about the client's first address, which is the client's
|
// GeoJS is asked about the client's first address, which is the client's
|
||||||
// own address for IPv4, and an address in the same place for an IPv6 /64.
|
// own address for IPv4, and an address in the same place for an IPv6 /64.
|
||||||
func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
|
||||||
country, asked := g.answerOrWait(ctx, client)
|
answer, asked := g.answerOrWait(ctx, client)
|
||||||
if asked == nil {
|
if asked == nil {
|
||||||
return country
|
return answer
|
||||||
}
|
}
|
||||||
|
|
||||||
timer := time.NewTimer(timeout)
|
timer := time.NewTimer(g.timeout)
|
||||||
defer timer.Stop()
|
defer timer.Stop()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
@@ -165,7 +195,7 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
|||||||
g.mu.Lock()
|
g.mu.Lock()
|
||||||
defer g.mu.Unlock()
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
country, found := g.kept(client)
|
answer, found := g.kept(client)
|
||||||
if !found {
|
if !found {
|
||||||
g.metrics.GeoJSUnanswered.Inc()
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
}
|
}
|
||||||
@@ -175,7 +205,15 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
|||||||
w.late = true
|
w.late = true
|
||||||
}
|
}
|
||||||
|
|
||||||
return country
|
return answer
|
||||||
|
}
|
||||||
|
|
||||||
|
// Kept returns client's answer, if one is kept, without asking GeoJS.
|
||||||
|
func (g *GeoJS) Kept(client netip.Prefix) (Answer, bool) {
|
||||||
|
g.mu.Lock()
|
||||||
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
|
return g.kept(client)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Snapshot returns every answer kept, sorted by client, as lookups.json
|
// Snapshot returns every answer kept, sorted by client, as lookups.json
|
||||||
@@ -227,13 +265,13 @@ func (g *GeoJS) Load(answers []Answer) {
|
|||||||
// nil when there is nothing to wait for.
|
// nil when there is nothing to wait for.
|
||||||
func (g *GeoJS) answerOrWait(
|
func (g *GeoJS) answerOrWait(
|
||||||
ctx context.Context, client netip.Prefix,
|
ctx context.Context, client netip.Prefix,
|
||||||
) (string, <-chan struct{}) {
|
) (Answer, <-chan struct{}) {
|
||||||
g.mu.Lock()
|
g.mu.Lock()
|
||||||
defer g.mu.Unlock()
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
country, found := g.kept(client)
|
answer, found := g.kept(client)
|
||||||
if found {
|
if found {
|
||||||
return country, nil
|
return answer, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
w, waiting := g.waiting[client]
|
w, waiting := g.waiting[client]
|
||||||
@@ -244,10 +282,14 @@ func (g *GeoJS) answerOrWait(
|
|||||||
|
|
||||||
g.ask(ctx)
|
g.ask(ctx)
|
||||||
|
|
||||||
|
if !g.wait {
|
||||||
|
return Answer{}, nil // the answer is not needed before the request goes on
|
||||||
|
}
|
||||||
|
|
||||||
if w == nil {
|
if w == nil {
|
||||||
g.metrics.GeoJSUnanswered.Inc()
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
|
|
||||||
return "", nil // too many clients wait already
|
return Answer{}, nil // too many clients wait already
|
||||||
}
|
}
|
||||||
|
|
||||||
if !g.asking {
|
if !g.asking {
|
||||||
@@ -258,25 +300,25 @@ func (g *GeoJS) answerOrWait(
|
|||||||
if w.late {
|
if w.late {
|
||||||
g.metrics.GeoJSUnanswered.Inc()
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
|
|
||||||
return "", nil
|
return Answer{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return "", w.asked
|
return Answer{}, w.asked
|
||||||
}
|
}
|
||||||
|
|
||||||
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
|
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
|
||||||
// and notes that it was used.
|
// and notes that it was used.
|
||||||
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
|
func (g *GeoJS) kept(client netip.Prefix) (Answer, bool) {
|
||||||
now := g.now()
|
now := g.now()
|
||||||
|
|
||||||
kept, found := g.answers.Get(client)
|
kept, found := g.answers.Get(client)
|
||||||
if !found || now.Sub(kept.Answered) >= keepFor {
|
if !found || now.Sub(kept.Answered) >= keepFor {
|
||||||
return "", false
|
return Answer{}, false
|
||||||
}
|
}
|
||||||
|
|
||||||
kept.Used = now
|
kept.Used = now
|
||||||
|
|
||||||
return kept.Country, true
|
return *kept, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// ask starts asking GeoJS about the waiting clients, unless a request to
|
// ask starts asking GeoJS about the waiting clients, unless a request to
|
||||||
@@ -294,7 +336,8 @@ func (g *GeoJS) ask(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// askAboutWaiting asks GeoJS about the waiting clients, one request at a
|
// askAboutWaiting asks GeoJS about the waiting clients, one request at a
|
||||||
// time, until none is left or GeoJS fails.
|
// time, until none is left or GeoJS fails. Each answer kept is given to
|
||||||
|
// Answered, outside the lock, since Answered takes locks of its own.
|
||||||
func (g *GeoJS) askAboutWaiting(ctx context.Context) {
|
func (g *GeoJS) askAboutWaiting(ctx context.Context) {
|
||||||
for {
|
for {
|
||||||
clients := g.nextClients()
|
clients := g.nextClients()
|
||||||
@@ -302,8 +345,16 @@ func (g *GeoJS) askAboutWaiting(ctx context.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
countries, err := g.request(ctx, clients)
|
given, err := g.request(ctx, clients)
|
||||||
if !g.keep(clients, countries, err) {
|
kept, answered := g.keep(clients, given, err)
|
||||||
|
|
||||||
|
if g.answered != nil {
|
||||||
|
for _, answer := range kept {
|
||||||
|
g.answered(answer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !answered {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -335,32 +386,35 @@ func (g *GeoJS) nextClients() []netip.Prefix {
|
|||||||
return clients
|
return clients
|
||||||
}
|
}
|
||||||
|
|
||||||
// keep notes how a request to GeoJS about clients ended, and reports
|
// keep notes how a request to GeoJS about clients ended, given being the
|
||||||
// whether GeoJS answered about all of them. Each client whose address
|
// answer for each address GeoJS's answer names. It returns the answers it
|
||||||
// GeoJS's answer names gets its answer, with no country when GeoJS gave
|
// kept, and reports whether GeoJS answered about all of the clients. Each
|
||||||
// none. An answer that leaves an address out is a failure. After a
|
// client whose address GeoJS's answer names gets its answer. An answer
|
||||||
// failure GeoJS is left alone for a while, and every client still waiting
|
// that leaves an address out is a failure. After a failure GeoJS is left
|
||||||
// stops waiting and is asked about once GeoJS is asked again.
|
// alone for a while, and every client still waiting stops waiting and is
|
||||||
|
// asked about once GeoJS is asked again.
|
||||||
func (g *GeoJS) keep(
|
func (g *GeoJS) keep(
|
||||||
clients []netip.Prefix, countries map[netip.Addr]string, err error,
|
clients []netip.Prefix, given map[netip.Addr]Answer, err error,
|
||||||
) bool {
|
) ([]Answer, bool) {
|
||||||
g.mu.Lock()
|
g.mu.Lock()
|
||||||
defer g.mu.Unlock()
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
now := g.now()
|
now := g.now()
|
||||||
|
kept := make([]Answer, 0, len(clients))
|
||||||
leftOut := 0
|
leftOut := 0
|
||||||
|
|
||||||
for _, client := range clients {
|
for _, client := range clients {
|
||||||
country, named := countries[client.Addr()]
|
answer, named := given[client.Addr()]
|
||||||
if !named {
|
if !named {
|
||||||
leftOut++
|
leftOut++
|
||||||
|
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
g.answers.Add(client, &Answer{
|
answer.Client, answer.Answered, answer.Used = client, now, now
|
||||||
Client: client, Country: country, Answered: now, Used: now,
|
g.answers.Add(client, &answer)
|
||||||
})
|
kept = append(kept, answer)
|
||||||
|
|
||||||
close(g.waiting[client].asked)
|
close(g.waiting[client].asked)
|
||||||
delete(g.waiting, client)
|
delete(g.waiting, client)
|
||||||
}
|
}
|
||||||
@@ -386,27 +440,37 @@ func (g *GeoJS) keep(
|
|||||||
|
|
||||||
g.processLog.Warn("asking GeoJS failed",
|
g.processLog.Warn("asking GeoJS failed",
|
||||||
"error", err.Error(), "asking_again_in", g.retryDelay.String())
|
"error", err.Error(), "asking_again_in", g.retryDelay.String())
|
||||||
|
g.alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventSourceFailure,
|
||||||
|
Reason: "asking GeoJS failed",
|
||||||
|
Detail: map[string]any{
|
||||||
|
"source": "geojs", "error": err.Error(),
|
||||||
|
"asking_again_in": g.retryDelay.String(),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
return false
|
return kept, false
|
||||||
}
|
}
|
||||||
|
|
||||||
g.retryDelay = 0
|
g.retryDelay = 0
|
||||||
|
|
||||||
return true
|
return kept, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// request asks GeoJS about clients in one request, and returns the
|
// request asks GeoJS about clients in one request, and returns the answer
|
||||||
// country it gave, in capitals, for each address its answer names.
|
// for each address GeoJS's answer names: its AS number and the AS's name,
|
||||||
|
// both "" for the AS number 64512, which GeoJS gives when it knows none,
|
||||||
|
// and its country, in capitals.
|
||||||
func (g *GeoJS) request(
|
func (g *GeoJS) request(
|
||||||
ctx context.Context, clients []netip.Prefix,
|
ctx context.Context, clients []netip.Prefix,
|
||||||
) (map[netip.Addr]string, error) {
|
) (map[netip.Addr]Answer, error) {
|
||||||
addrs := make([]string, 0, len(clients))
|
addrs := make([]string, 0, len(clients))
|
||||||
|
|
||||||
for _, client := range clients {
|
for _, client := range clients {
|
||||||
addrs = append(addrs, client.Addr().String())
|
addrs = append(addrs, client.Addr().String())
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
ctx, cancel := context.WithTimeout(ctx, g.timeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody)
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody)
|
||||||
@@ -433,9 +497,12 @@ func (g *GeoJS) request(
|
|||||||
return nil, fmt.Errorf("%w %s", errStatus, res.Status)
|
return nil, fmt.Errorf("%w %s", errStatus, res.Status)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
//nolint:tagliatelle // GeoJS's own names
|
||||||
var answers []struct {
|
var answers []struct {
|
||||||
IP string `json:"ip"`
|
IP string `json:"ip"`
|
||||||
Country string `json:"country"`
|
ASN int64 `json:"asn"`
|
||||||
|
ASName string `json:"organization_name"`
|
||||||
|
CountryCode string `json:"country_code"`
|
||||||
}
|
}
|
||||||
|
|
||||||
err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers)
|
err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers)
|
||||||
@@ -443,14 +510,22 @@ func (g *GeoJS) request(
|
|||||||
return nil, fmt.Errorf("read GeoJS's answer: %w", err)
|
return nil, fmt.Errorf("read GeoJS's answer: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
countries := make(map[netip.Addr]string, len(answers))
|
given := make(map[netip.Addr]Answer, len(answers))
|
||||||
|
|
||||||
for _, item := range answers {
|
for _, item := range answers {
|
||||||
addr, err := netip.ParseAddr(item.IP)
|
addr, err := netip.ParseAddr(item.IP)
|
||||||
if err == nil {
|
if err != nil {
|
||||||
countries[addr] = strings.ToUpper(item.Country)
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
answer := Answer{Country: strings.ToUpper(item.CountryCode)}
|
||||||
|
if item.ASN != 0 && item.ASN != unknownASN {
|
||||||
|
answer.ASN = "AS" + strconv.FormatInt(item.ASN, 10)
|
||||||
|
answer.ASName = item.ASName
|
||||||
|
}
|
||||||
|
|
||||||
|
given[addr] = answer
|
||||||
}
|
}
|
||||||
|
|
||||||
return countries, nil
|
return given, nil
|
||||||
}
|
}
|
||||||
|
|||||||
+247
-16
@@ -6,6 +6,8 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"net/url"
|
||||||
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -14,15 +16,20 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus/testutil"
|
"github.com/prometheus/client_golang/prometheus/testutil"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// germany is where the stand-in for GeoJS places every address but
|
// germany is where the stand-in for GeoJS places every address but
|
||||||
// unplaced.
|
// unplaced, and asNumber, kept as asn, and asName the AS it gives them.
|
||||||
germany = "DE"
|
germany = "DE"
|
||||||
// unplaced is the address it cannot place.
|
asNumber = 64496
|
||||||
|
asn = "AS64496"
|
||||||
|
asName = "Example Net"
|
||||||
|
// unplaced is the address it cannot place, for which it gives the AS
|
||||||
|
// number 64512 and the AS name Unknown, as GeoJS does.
|
||||||
unplaced = "192.0.2.1"
|
unplaced = "192.0.2.1"
|
||||||
// leftOut is the address it leaves out of its answer when
|
// leftOut is the address it leaves out of its answer when
|
||||||
// answeringWithoutLeftOut.
|
// answeringWithoutLeftOut.
|
||||||
@@ -80,7 +87,7 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
|
|||||||
|
|
||||||
var earlier sync.WaitGroup
|
var earlier sync.WaitGroup
|
||||||
|
|
||||||
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
|
earlier.Go(func() { g.LookUp(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
|
||||||
defer earlier.Wait()
|
defer earlier.Wait()
|
||||||
|
|
||||||
waitForRequests(t, geojs, 1)
|
waitForRequests(t, geojs, 1)
|
||||||
@@ -112,6 +119,58 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRequestWaitsAsLongAsTheTimeoutSaysAndGeoJSIsAbandonedAfterIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
// A timeout longer than the default second, and a GeoJS that does
|
||||||
|
// not answer.
|
||||||
|
const longerTimeout = 3 * time.Second
|
||||||
|
|
||||||
|
m := metrics.New(1, "app")
|
||||||
|
g := lookup.New(lookup.Params{
|
||||||
|
URL: lookup.URL,
|
||||||
|
Timeout: longerTimeout,
|
||||||
|
Wait: true,
|
||||||
|
Now: time.Now,
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Metrics: m,
|
||||||
|
Alerts: alerts.New(alerts.Params{}),
|
||||||
|
})
|
||||||
|
g.SetTransport(&standIn{answers: hanging})
|
||||||
|
|
||||||
|
var (
|
||||||
|
request sync.WaitGroup
|
||||||
|
waited time.Duration
|
||||||
|
)
|
||||||
|
|
||||||
|
request.Go(func() {
|
||||||
|
began := time.Now()
|
||||||
|
|
||||||
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
||||||
|
|
||||||
|
waited = time.Since(began)
|
||||||
|
})
|
||||||
|
|
||||||
|
// A moment before the timeout runs out, GeoJS is still being asked:
|
||||||
|
// the request to it has not failed.
|
||||||
|
time.Sleep(longerTimeout - time.Millisecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFailures(t, m, 0)
|
||||||
|
|
||||||
|
// As it runs out, the client's request goes on, and the request to
|
||||||
|
// GeoJS is abandoned, which counts as a failure.
|
||||||
|
request.Wait()
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
if waited != longerTimeout {
|
||||||
|
t.Errorf("waited %s for the answer, want %s", waited, longerTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantFailures(t, m, 1)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -184,6 +243,96 @@ func TestCountryIsKeptInCapitals(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAnswerHoldsTheASNumberTheASNameAndTheCountry(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
_, clock, g := start()
|
||||||
|
placed := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
notPlaced := netip.MustParsePrefix(unplaced + "/32")
|
||||||
|
now := clock.Now()
|
||||||
|
|
||||||
|
// For the client it cannot place, GeoJS gives the AS number 64512
|
||||||
|
// and the AS name Unknown, which count as unknown.
|
||||||
|
for client, want := range map[netip.Prefix]lookup.Answer{
|
||||||
|
placed: {
|
||||||
|
Client: placed, ASN: asn, ASName: asName, Country: germany,
|
||||||
|
Answered: now, Used: now,
|
||||||
|
},
|
||||||
|
notPlaced: {Client: notPlaced, Answered: now, Used: now},
|
||||||
|
} {
|
||||||
|
got := g.LookUp(t.Context(), client)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("answer for %s\n%+v\nwant\n%+v", client, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWithoutWaitTheRequestGoesOnAtOnceAndTheAnswerIsGivenWhenItComes(
|
||||||
|
t *testing.T,
|
||||||
|
) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
var (
|
||||||
|
mu sync.Mutex
|
||||||
|
given []lookup.Answer
|
||||||
|
)
|
||||||
|
|
||||||
|
geojs := &standIn{answers: answeringSlowly}
|
||||||
|
clock := newClock()
|
||||||
|
m := metrics.New(1, "app")
|
||||||
|
g := lookup.New(lookup.Params{
|
||||||
|
URL: lookup.URL,
|
||||||
|
Timeout: timeout,
|
||||||
|
Answered: func(answer lookup.Answer) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
given = append(given, answer)
|
||||||
|
},
|
||||||
|
Now: clock.Now,
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Metrics: m,
|
||||||
|
Alerts: alerts.New(alerts.Params{}),
|
||||||
|
})
|
||||||
|
g.SetTransport(geojs)
|
||||||
|
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
// The request goes on at once, without an answer, and GeoJS is asked
|
||||||
|
// about the client, which it answers most of a second later.
|
||||||
|
began := time.Now()
|
||||||
|
|
||||||
|
got := g.LookUp(t.Context(), client)
|
||||||
|
if took := time.Since(began); took != 0 || got != (lookup.Answer{}) {
|
||||||
|
t.Errorf("waited %s for %+v, want no wait and no answer", took, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
waitForRequests(t, geojs, 1)
|
||||||
|
wantAsked(t, geojs, 0, "203.0.113.9")
|
||||||
|
|
||||||
|
time.Sleep(timeout)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
now := clock.Now()
|
||||||
|
want := lookup.Answer{
|
||||||
|
Client: client, ASN: asn, ASName: asName, Country: germany,
|
||||||
|
Answered: now, Used: now,
|
||||||
|
}
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
if !slices.Equal(given, []lookup.Answer{want}) {
|
||||||
|
t.Errorf("answers given %+v, want only %+v", given, want)
|
||||||
|
}
|
||||||
|
mu.Unlock()
|
||||||
|
|
||||||
|
wantCountry(t, g, client, germany)
|
||||||
|
wantUnanswered(t, m, 0)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -194,9 +343,12 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
|||||||
geojs := &standIn{answers: hanging}
|
geojs := &standIn{answers: hanging}
|
||||||
g := lookup.New(lookup.Params{
|
g := lookup.New(lookup.Params{
|
||||||
URL: lookup.URL,
|
URL: lookup.URL,
|
||||||
|
Timeout: timeout,
|
||||||
|
Wait: true,
|
||||||
Now: time.Now,
|
Now: time.Now,
|
||||||
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
|
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
|
||||||
Metrics: metrics.New(1),
|
Metrics: metrics.New(1, "app"),
|
||||||
|
Alerts: alerts.New(alerts.Params{}),
|
||||||
})
|
})
|
||||||
g.SetTransport(geojs)
|
g.SetTransport(geojs)
|
||||||
|
|
||||||
@@ -211,6 +363,45 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
geojs, clock, g, queue := startWithAlerts()
|
||||||
|
clients := newClients()
|
||||||
|
|
||||||
|
geojs.set(failing)
|
||||||
|
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
|
||||||
|
want := alerts.Alert{
|
||||||
|
Time: clock.Now(),
|
||||||
|
Event: alerts.EventSourceFailure,
|
||||||
|
Reason: "asking GeoJS failed",
|
||||||
|
Detail: map[string]any{
|
||||||
|
"source": "geojs",
|
||||||
|
"error": "GeoJS answered 503 Service Unavailable",
|
||||||
|
"asking_again_in": "1s",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next failure, a second later, is a repeat within the
|
||||||
|
// cooldown.
|
||||||
|
clock.advance(time.Second)
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
wantRequests(t, geojs, 2)
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
|
||||||
|
t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if queue.Suppressed() != 1 {
|
||||||
|
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
|
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -355,12 +546,15 @@ func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
synctest.Test(t, func(t *testing.T) {
|
||||||
m := metrics.New(1)
|
m := metrics.New(1, "app")
|
||||||
g := lookup.New(lookup.Params{
|
g := lookup.New(lookup.Params{
|
||||||
URL: lookup.URL,
|
URL: lookup.URL,
|
||||||
|
Timeout: timeout,
|
||||||
|
Wait: true,
|
||||||
Now: time.Now,
|
Now: time.Now,
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
Metrics: m,
|
Metrics: m,
|
||||||
|
Alerts: alerts.New(alerts.Params{}),
|
||||||
})
|
})
|
||||||
g.SetTransport(&standIn{answers: failing})
|
g.SetTransport(&standIn{answers: failing})
|
||||||
|
|
||||||
@@ -450,21 +644,24 @@ func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
list := make([]map[string]string, 0, len(addrs))
|
list := make([]map[string]any, 0, len(addrs))
|
||||||
|
|
||||||
for _, addr := range addrs {
|
for _, addr := range addrs {
|
||||||
country := germany
|
item := map[string]any{
|
||||||
|
"ip": addr, "asn": asNumber, "organization_name": asName,
|
||||||
|
"country_code": germany,
|
||||||
|
}
|
||||||
|
|
||||||
switch {
|
switch {
|
||||||
case addr == unplaced:
|
case addr == unplaced:
|
||||||
country = ""
|
item = map[string]any{"ip": addr, "asn": 64512, "organization_name": "Unknown"}
|
||||||
case addr == leftOut && answers == answeringWithoutLeftOut:
|
case addr == leftOut && answers == answeringWithoutLeftOut:
|
||||||
continue
|
continue
|
||||||
case answers == answeringInLowerCase:
|
case answers == answeringInLowerCase:
|
||||||
country = strings.ToLower(germany)
|
item["country_code"] = strings.ToLower(germany)
|
||||||
}
|
}
|
||||||
|
|
||||||
list = append(list, map[string]string{"ip": addr, "country": country})
|
list = append(list, item)
|
||||||
}
|
}
|
||||||
|
|
||||||
var answer any = list
|
var answer any = list
|
||||||
@@ -522,19 +719,43 @@ func (c *testClock) advance(d time.Duration) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
|
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
|
||||||
// asking the stand-in by that clock.
|
// asking the stand-in by that clock, for which a request waits for its
|
||||||
|
// client's first answer.
|
||||||
func start() (*standIn, *testClock, *lookup.GeoJS) {
|
func start() (*standIn, *testClock, *lookup.GeoJS) {
|
||||||
|
geojs, clock, g, _ := startWithAlerts()
|
||||||
|
|
||||||
|
return geojs, clock, g
|
||||||
|
}
|
||||||
|
|
||||||
|
// startWithAlerts is start, and returns the queue of the alerts GeoJS
|
||||||
|
// raises as well, for a webhook that is never sent them, with the default
|
||||||
|
// cooldown, by the same clock.
|
||||||
|
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
|
||||||
geojs := &standIn{}
|
geojs := &standIn{}
|
||||||
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
|
clock := newClock()
|
||||||
|
queue := alerts.New(alerts.Params{
|
||||||
|
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||||
|
Events: alerts.Events(),
|
||||||
|
Cooldown: 15 * time.Minute,
|
||||||
|
Now: clock.Now,
|
||||||
|
})
|
||||||
g := lookup.New(lookup.Params{
|
g := lookup.New(lookup.Params{
|
||||||
URL: lookup.URL,
|
URL: lookup.URL,
|
||||||
|
Timeout: timeout,
|
||||||
|
Wait: true,
|
||||||
Now: clock.Now,
|
Now: clock.Now,
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
Metrics: metrics.New(1),
|
Metrics: metrics.New(1, "app"),
|
||||||
|
Alerts: queue,
|
||||||
})
|
})
|
||||||
g.SetTransport(geojs)
|
g.SetTransport(geojs)
|
||||||
|
|
||||||
return geojs, clock, g
|
return geojs, clock, g, queue
|
||||||
|
}
|
||||||
|
|
||||||
|
// newClock returns a clock set to the start of a day.
|
||||||
|
func newClock() *testClock {
|
||||||
|
return &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
|
||||||
}
|
}
|
||||||
|
|
||||||
// newClients returns what returns a new IPv4 client each time it is
|
// newClients returns what returns a new IPv4 client each time it is
|
||||||
@@ -553,7 +774,7 @@ func newClients() func() netip.Prefix {
|
|||||||
func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) {
|
func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
got := g.Country(t.Context(), client)
|
got := g.LookUp(t.Context(), client).Country
|
||||||
if got != want {
|
if got != want {
|
||||||
t.Errorf("%s is in %q, want %q", client, got, want)
|
t.Errorf("%s is in %q, want %q", client, got, want)
|
||||||
}
|
}
|
||||||
@@ -598,6 +819,16 @@ func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// wantFailures checks how many requests to GeoJS m counts as failed.
|
||||||
|
func wantFailures(t *testing.T, m *metrics.Metrics, want float64) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := testutil.ToFloat64(m.GeoJSFailures)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("%v requests to GeoJS failed, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// waitForRequests waits until g has done all it can before time passes,
|
// waitForRequests waits until g has done all it can before time passes,
|
||||||
// checks that GeoJS has had count requests, and returns the addresses each
|
// checks that GeoJS has had count requests, and returns the addresses each
|
||||||
// asked about.
|
// asked about.
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
// Package lookuptest writes lookup databases, IPinfo Lite files in their
|
||||||
|
// .mmdb form, for the tests of the packages that read them.
|
||||||
|
package lookuptest
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/maxmind/mmdbwriter"
|
||||||
|
"github.com/maxmind/mmdbwriter/mmdbtype"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fileMode is the mode of the files written: read and written by their
|
||||||
|
// owner alone.
|
||||||
|
const fileMode = 0o600
|
||||||
|
|
||||||
|
// Network is what a lookup database holds about a netblock, of the fields
|
||||||
|
// smallwebwaf reads: its AS number, such as AS64496, the AS's name, and
|
||||||
|
// its country, such as DE.
|
||||||
|
type Network struct {
|
||||||
|
ASN string
|
||||||
|
ASName string
|
||||||
|
Country string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write writes a lookup database at path that places each netblock in
|
||||||
|
// networks, such as 203.0.113.0/24, as its Network says, and no other
|
||||||
|
// address.
|
||||||
|
func Write(tb testing.TB, path string, networks map[string]Network) {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
records := make(map[string]mmdbtype.Map, len(networks))
|
||||||
|
for netblock, network := range networks {
|
||||||
|
records[netblock] = mmdbtype.Map{
|
||||||
|
"asn": mmdbtype.String(network.ASN),
|
||||||
|
"as_name": mmdbtype.String(network.ASName),
|
||||||
|
"country_code": mmdbtype.String(network.Country),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
WriteRecords(tb, path, records)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteRecords writes a lookup database at path that holds each record in
|
||||||
|
// records for its netblock, and nothing for any other address.
|
||||||
|
func WriteRecords(tb testing.TB, path string, records map[string]mmdbtype.Map) {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
tree, err := mmdbwriter.New(mmdbwriter.Options{
|
||||||
|
DatabaseType: "ipinfo_lite",
|
||||||
|
// The tests' clients are in the netblocks kept for documentation.
|
||||||
|
IncludeReservedNetworks: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
tb.Fatalf("new lookup database: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for netblock, record := range records {
|
||||||
|
_, network, err := net.ParseCIDR(netblock)
|
||||||
|
if err != nil {
|
||||||
|
tb.Fatalf("netblock %q: %v", netblock, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = tree.Insert(network, record)
|
||||||
|
if err != nil {
|
||||||
|
tb.Fatalf("insert %s: %v", netblock, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var database bytes.Buffer
|
||||||
|
|
||||||
|
_, err = tree.WriteTo(&database)
|
||||||
|
if err != nil {
|
||||||
|
tb.Fatalf("write the lookup database: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = os.WriteFile(path, database.Bytes(), fileMode)
|
||||||
|
if err != nil {
|
||||||
|
tb.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -26,8 +26,11 @@ func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
|
|||||||
wantCountry(t, g, placed, germany)
|
wantCountry(t, g, placed, germany)
|
||||||
|
|
||||||
want := []lookup.Answer{
|
want := []lookup.Answer{
|
||||||
{Client: notPlaced, Country: "", Answered: asked, Used: asked},
|
{Client: notPlaced, Answered: asked, Used: asked},
|
||||||
{Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
|
{
|
||||||
|
Client: placed, ASN: asn, ASName: asName, Country: germany,
|
||||||
|
Answered: asked, Used: asked.Add(time.Hour),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
if got := g.Snapshot(); !slices.Equal(got, want) {
|
if got := g.Snapshot(); !slices.Equal(got, want) {
|
||||||
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
|
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
|
||||||
|
|||||||
@@ -0,0 +1,155 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// other is the label value under which the countries or AS numbers
|
||||||
|
// outside the busiest are counted.
|
||||||
|
const other = "other"
|
||||||
|
|
||||||
|
// busiest are the metrics by one thing the lookup finds of the client,
|
||||||
|
// its country or its AS number, for requests whose client's is known.
|
||||||
|
// The topN busiest countries or AS numbers, by their requests since the
|
||||||
|
// start, have series of their own, and the others are counted under
|
||||||
|
// other, so that there are never more than topN + 1 series. One that
|
||||||
|
// drops out of the busiest loses its series, and its next requests are
|
||||||
|
// counted under other; one that becomes one of them gets a series that
|
||||||
|
// counts from then on. Each series therefore only ever goes up.
|
||||||
|
type busiest struct {
|
||||||
|
topN int
|
||||||
|
|
||||||
|
requests *prometheus.CounterVec
|
||||||
|
requestBytes *prometheus.CounterVec
|
||||||
|
responseBytes *prometheus.CounterVec
|
||||||
|
// refused are, by country, the requests the country lists refused; nil
|
||||||
|
// by AS number.
|
||||||
|
refused *prometheus.CounterVec
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
// seen is each country's or AS number's requests since the start, by
|
||||||
|
// which they are ranked.
|
||||||
|
seen map[string]int64
|
||||||
|
// top are those with series of their own.
|
||||||
|
top map[string]bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// newCountries returns the metrics by country, with series of their own
|
||||||
|
// for the topN busiest countries.
|
||||||
|
func newCountries(topN int) *busiest {
|
||||||
|
countries := newBusiest(topN, "country", "the client's country")
|
||||||
|
countries.refused = counterVec("smallwebwaf_country_list_refusals_total",
|
||||||
|
"Requests the country lists refused, by the client's country.",
|
||||||
|
[]string{"country"})
|
||||||
|
|
||||||
|
return countries
|
||||||
|
}
|
||||||
|
|
||||||
|
// newASNs returns the metrics by AS number, with series of their own for
|
||||||
|
// the topN busiest AS numbers.
|
||||||
|
func newASNs(topN int) *busiest {
|
||||||
|
return newBusiest(topN, "asn", "the client's AS number")
|
||||||
|
}
|
||||||
|
|
||||||
|
// newBusiest returns the metrics by label, which is described as
|
||||||
|
// description, with series of their own for the topN busiest values.
|
||||||
|
func newBusiest(topN int, label, description string) *busiest {
|
||||||
|
by := []string{label}
|
||||||
|
|
||||||
|
return &busiest{
|
||||||
|
topN: topN,
|
||||||
|
requests: counterVec("smallwebwaf_"+label+"_requests_total",
|
||||||
|
"Requests, by "+description+".", by),
|
||||||
|
requestBytes: counterVec("smallwebwaf_"+label+"_request_bytes_total",
|
||||||
|
"Request body bytes, by "+description+".", by),
|
||||||
|
responseBytes: counterVec("smallwebwaf_"+label+"_response_bytes_total",
|
||||||
|
"Response body bytes, by "+description+".", by),
|
||||||
|
seen: map[string]int64{},
|
||||||
|
top: map[string]bool{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Describe and Collect make the metrics a prometheus.Collector, so that
|
||||||
|
// they are registered together.
|
||||||
|
func (b *busiest) Describe(ch chan<- *prometheus.Desc) {
|
||||||
|
for _, vec := range b.vecs() {
|
||||||
|
vec.Describe(ch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Collect is the other half of prometheus.Collector, with Describe.
|
||||||
|
func (b *busiest) Collect(ch chan<- prometheus.Metric) {
|
||||||
|
for _, vec := range b.vecs() {
|
||||||
|
vec.Collect(ch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// vecs returns the metrics: by AS number, those of requests and bytes; by
|
||||||
|
// country, the refusals by the country lists as well.
|
||||||
|
func (b *busiest) vecs() []*prometheus.CounterVec {
|
||||||
|
vecs := []*prometheus.CounterVec{b.requests, b.requestBytes, b.responseBytes}
|
||||||
|
if b.refused != nil {
|
||||||
|
vecs = append(vecs, b.refused)
|
||||||
|
}
|
||||||
|
|
||||||
|
return vecs
|
||||||
|
}
|
||||||
|
|
||||||
|
// add counts a request from its log line, whose client's country or AS
|
||||||
|
// number, value, is known.
|
||||||
|
func (b *busiest) add(value string, line *requestlog.Line) {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
|
||||||
|
b.seen[value]++
|
||||||
|
|
||||||
|
label := b.label(value)
|
||||||
|
b.requests.WithLabelValues(label).Inc()
|
||||||
|
b.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
|
||||||
|
b.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
|
||||||
|
|
||||||
|
if b.refused != nil && line.Action == requestlog.ActionCountryDenied {
|
||||||
|
b.refused.WithLabelValues(label).Inc()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// label returns the label a request from value is counted under: value
|
||||||
|
// while it is one of the busiest, other while it is not. A value busier
|
||||||
|
// than the least busy of them takes its place, and that one's series are
|
||||||
|
// dropped.
|
||||||
|
func (b *busiest) label(value string) string {
|
||||||
|
if b.top[value] {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(b.top) < b.topN {
|
||||||
|
b.top[value] = true
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
least := ""
|
||||||
|
|
||||||
|
for top := range b.top {
|
||||||
|
if least == "" || b.seen[top] < b.seen[least] {
|
||||||
|
least = top
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if b.seen[value] <= b.seen[least] {
|
||||||
|
return other
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(b.top, least)
|
||||||
|
|
||||||
|
for _, vec := range b.vecs() {
|
||||||
|
vec.DeleteLabelValues(least)
|
||||||
|
}
|
||||||
|
|
||||||
|
b.top[value] = true
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
@@ -1,116 +0,0 @@
|
|||||||
package metrics
|
|
||||||
|
|
||||||
import (
|
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// other is the label under which the countries outside the busiest are
|
|
||||||
// counted.
|
|
||||||
const other = "other"
|
|
||||||
|
|
||||||
// countries are the metrics by the client's country, for requests whose
|
|
||||||
// client's country is known. The topN busiest countries, by their requests
|
|
||||||
// since the start, have series of their own, and the others are counted
|
|
||||||
// under other, so that there are never more than topN + 1 series. A
|
|
||||||
// country that drops out of the busiest loses its series, and its next
|
|
||||||
// requests are counted under other; one that becomes one of them gets a
|
|
||||||
// series that counts from then on. Each series therefore only ever goes
|
|
||||||
// up.
|
|
||||||
type countries struct {
|
|
||||||
topN int
|
|
||||||
|
|
||||||
requests *prometheus.CounterVec
|
|
||||||
requestBytes *prometheus.CounterVec
|
|
||||||
responseBytes *prometheus.CounterVec
|
|
||||||
// refused are the requests the country lists refused.
|
|
||||||
refused *prometheus.CounterVec
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// seen is each country's requests since the start, by which the
|
|
||||||
// countries are ranked. GeoJS gives two-letter codes, so it holds at
|
|
||||||
// most a few hundred.
|
|
||||||
seen map[string]int64
|
|
||||||
// top are the countries with series of their own.
|
|
||||||
top map[string]bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// newCountries returns the metrics by country, with series of their own
|
|
||||||
// for the topN busiest countries.
|
|
||||||
func newCountries(topN int) *countries {
|
|
||||||
byCountry := []string{"country"}
|
|
||||||
|
|
||||||
return &countries{
|
|
||||||
topN: topN,
|
|
||||||
requests: counterVec("smallwebwaf_country_requests_total",
|
|
||||||
"Requests, by the client's country.", byCountry),
|
|
||||||
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
|
|
||||||
"Request body bytes, by the client's country.", byCountry),
|
|
||||||
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
|
|
||||||
"Response body bytes, by the client's country.", byCountry),
|
|
||||||
refused: counterVec("smallwebwaf_country_list_refusals_total",
|
|
||||||
"Requests the country lists refused, by the client's country.",
|
|
||||||
byCountry),
|
|
||||||
seen: map[string]int64{},
|
|
||||||
top: map[string]bool{},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// add counts a request from its log line, whose country is known.
|
|
||||||
func (c *countries) add(line *requestlog.Line) {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
|
|
||||||
c.seen[line.Country]++
|
|
||||||
|
|
||||||
label := c.label(line.Country)
|
|
||||||
c.requests.WithLabelValues(label).Inc()
|
|
||||||
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
|
|
||||||
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
|
|
||||||
|
|
||||||
if line.Action == requestlog.ActionCountryDenied {
|
|
||||||
c.refused.WithLabelValues(label).Inc()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// label returns the label a request from country is counted under: the
|
|
||||||
// country while it is one of the busiest, other while it is not. A
|
|
||||||
// country busier than the least busy of them takes its place, and that
|
|
||||||
// country's series are dropped.
|
|
||||||
func (c *countries) label(country string) string {
|
|
||||||
if c.top[country] {
|
|
||||||
return country
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(c.top) < c.topN {
|
|
||||||
c.top[country] = true
|
|
||||||
|
|
||||||
return country
|
|
||||||
}
|
|
||||||
|
|
||||||
least := ""
|
|
||||||
|
|
||||||
for top := range c.top {
|
|
||||||
if least == "" || c.seen[top] < c.seen[least] {
|
|
||||||
least = top
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if c.seen[country] <= c.seen[least] {
|
|
||||||
return other
|
|
||||||
}
|
|
||||||
|
|
||||||
delete(c.top, least)
|
|
||||||
|
|
||||||
for _, vec := range []*prometheus.CounterVec{
|
|
||||||
c.requests, c.requestBytes, c.responseBytes, c.refused,
|
|
||||||
} {
|
|
||||||
vec.DeleteLabelValues(least)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.top[country] = true
|
|
||||||
|
|
||||||
return country
|
|
||||||
}
|
|
||||||
+104
-20
@@ -6,11 +6,13 @@ package metrics
|
|||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
"github.com/prometheus/client_golang/prometheus/collectors"
|
"github.com/prometheus/client_golang/prometheus/collectors"
|
||||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
@@ -20,7 +22,8 @@ import (
|
|||||||
|
|
||||||
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
|
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
|
||||||
type Metrics struct {
|
type Metrics struct {
|
||||||
registry *prometheus.Registry
|
// registry gives every metric registered with it the label instance.
|
||||||
|
registry prometheus.Registerer
|
||||||
handler http.Handler
|
handler http.Handler
|
||||||
|
|
||||||
inFlight prometheus.Gauge
|
inFlight prometheus.Gauge
|
||||||
@@ -34,12 +37,13 @@ type Metrics struct {
|
|||||||
offences *prometheus.CounterVec
|
offences *prometheus.CounterVec
|
||||||
// ruleMatches are made by AddRules.
|
// ruleMatches are made by AddRules.
|
||||||
ruleMatches *prometheus.CounterVec
|
ruleMatches *prometheus.CounterVec
|
||||||
countries *countries
|
countries *busiest
|
||||||
|
asns *busiest
|
||||||
|
|
||||||
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
|
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
|
||||||
// that failed. GeoJSUnanswered are the requests whose client counted
|
// that failed. GeoJSUnanswered are the requests that needed their
|
||||||
// as coming from an unknown country because GeoJS had not answered
|
// client's answer, for a setting that acts on it, and went on without
|
||||||
// about it in time.
|
// it because GeoJS had not given it in time.
|
||||||
GeoJSRequests prometheus.Counter
|
GeoJSRequests prometheus.Counter
|
||||||
GeoJSFailures prometheus.Counter
|
GeoJSFailures prometheus.Counter
|
||||||
GeoJSUnanswered prometheus.Counter
|
GeoJSUnanswered prometheus.Counter
|
||||||
@@ -53,14 +57,18 @@ type Metrics struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// New returns the metrics, with the Go runtime's and the process's own.
|
// New returns the metrics, with the Go runtime's and the process's own.
|
||||||
// topN is how many countries get series of their own
|
// topN is how many countries and how many AS numbers get series of their
|
||||||
// (SWWAF_METRICS_TOP_N).
|
// own (SWWAF_METRICS_TOP_N). Every metric carries instanceName
|
||||||
func New(topN int) *Metrics {
|
// (SWWAF_INSTANCE_NAME) as its label instance.
|
||||||
|
func New(topN int, instanceName string) *Metrics {
|
||||||
byStatus := []string{"status_class", "action"}
|
byStatus := []string{"status_class", "action"}
|
||||||
byFile := []string{"file"}
|
byFile := []string{"file"}
|
||||||
|
registry := prometheus.NewRegistry()
|
||||||
|
|
||||||
m := &Metrics{
|
m := &Metrics{
|
||||||
registry: prometheus.NewRegistry(),
|
registry: prometheus.WrapRegistererWith(
|
||||||
|
prometheus.Labels{"instance": instanceName}, registry),
|
||||||
|
handler: promhttp.HandlerFor(registry, promhttp.HandlerOpts{}),
|
||||||
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
|
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
Name: "smallwebwaf_requests_in_flight",
|
Name: "smallwebwaf_requests_in_flight",
|
||||||
Help: "Requests under way.",
|
Help: "Requests under way.",
|
||||||
@@ -82,14 +90,16 @@ func New(topN int) *Metrics {
|
|||||||
Help: "How long requests passed to the app took, from then to their end.",
|
Help: "How long requests passed to the app took, from then to their end.",
|
||||||
}),
|
}),
|
||||||
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
|
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
|
||||||
"Requests that broke a rate limit, by its window.",
|
"Requests that broke a rate limit or a byte limit, by its window and "+
|
||||||
[]string{"window"}),
|
"its kind, requests or bytes.",
|
||||||
|
[]string{"window", "kind"}),
|
||||||
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
|
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
|
||||||
"Requests that passed a size or time limit, by its setting.",
|
"Requests that passed a size or time limit, by its setting.",
|
||||||
[]string{"limit"}),
|
[]string{"limit"}),
|
||||||
offences: counterVec("smallwebwaf_offences_total",
|
offences: counterVec("smallwebwaf_offences_total",
|
||||||
"Offences, by kind.", []string{"kind"}),
|
"Offences, by kind.", []string{"kind"}),
|
||||||
countries: newCountries(topN),
|
countries: newCountries(topN),
|
||||||
|
asns: newASNs(topN),
|
||||||
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
|
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
|
||||||
Name: "smallwebwaf_geojs_requests_total",
|
Name: "smallwebwaf_geojs_requests_total",
|
||||||
Help: "Requests to GeoJS.",
|
Help: "Requests to GeoJS.",
|
||||||
@@ -100,8 +110,8 @@ func New(topN int) *Metrics {
|
|||||||
}),
|
}),
|
||||||
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
|
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
|
||||||
Name: "smallwebwaf_geojs_unanswered_total",
|
Name: "smallwebwaf_geojs_unanswered_total",
|
||||||
Help: "Requests whose client counted as coming from an unknown " +
|
Help: "Requests that needed their client's answer from GeoJS and " +
|
||||||
"country because GeoJS had not answered about it in time.",
|
"went on without it, because GeoJS had not given it in time.",
|
||||||
}),
|
}),
|
||||||
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
|
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
|
||||||
"Writes of each state file.", byFile),
|
"Writes of each state file.", byFile),
|
||||||
@@ -118,16 +128,12 @@ func New(topN int) *Metrics {
|
|||||||
byFile),
|
byFile),
|
||||||
}
|
}
|
||||||
|
|
||||||
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
|
|
||||||
|
|
||||||
m.registry.MustRegister(
|
m.registry.MustRegister(
|
||||||
collectors.NewGoCollector(),
|
collectors.NewGoCollector(),
|
||||||
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
|
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
|
||||||
m.inFlight, m.requests, m.requestBytes, m.responseBytes,
|
m.inFlight, m.requests, m.requestBytes, m.responseBytes,
|
||||||
m.requestDuration, m.upstreamDuration,
|
m.requestDuration, m.upstreamDuration,
|
||||||
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
|
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.countries, m.asns,
|
||||||
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
|
|
||||||
m.countries.refused,
|
|
||||||
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
|
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
|
||||||
m.stateFileWrites, m.stateFileWriteFailures,
|
m.stateFileWrites, m.stateFileWriteFailures,
|
||||||
m.stateFileLastWrite, m.stateFileSize,
|
m.stateFileLastWrite, m.stateFileSize,
|
||||||
@@ -225,6 +231,72 @@ func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddLookupFile adds the metrics of the lookup database, read as the
|
||||||
|
// metrics are asked for: when the file in use was read, which lastRead
|
||||||
|
// returns, and the replacements of it that could not be read, which
|
||||||
|
// readFailures returns. The lookup package's File, which has both, cannot
|
||||||
|
// be named here: that package counts GeoJS's requests in these metrics.
|
||||||
|
func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() int) {
|
||||||
|
m.registry.MustRegister(
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_lookup_database_last_read_timestamp_seconds",
|
||||||
|
Help: "When the lookup database in use was read, in seconds since 1970.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(lastRead().Unix())
|
||||||
|
}),
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_lookup_database_read_failures_total",
|
||||||
|
Help: "Replacements of the lookup database that could not be read.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(readFailures())
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddAlerts adds the metrics of the alerts sent to each destination set,
|
||||||
|
// read from queue as the metrics are asked for, by destination: the
|
||||||
|
// alerts sent, the requests to the destination that failed, the alerts
|
||||||
|
// held back, which are the same for every destination, and those
|
||||||
|
// dropped. With no destination set, it adds none.
|
||||||
|
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
|
||||||
|
for _, name := range queue.DestinationsSet() {
|
||||||
|
destination := prometheus.Labels{"destination": name}
|
||||||
|
|
||||||
|
m.registry.MustRegister(
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_alerts_sent_total",
|
||||||
|
Help: "Alerts the destination took.",
|
||||||
|
ConstLabels: destination,
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(queue.Counts(name).Sent)
|
||||||
|
}),
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_alerts_failed_total",
|
||||||
|
Help: "Requests to the destination that failed.",
|
||||||
|
ConstLabels: destination,
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(queue.Counts(name).Failed)
|
||||||
|
}),
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_alerts_suppressed_total",
|
||||||
|
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
|
||||||
|
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
|
||||||
|
ConstLabels: destination,
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(queue.Suppressed())
|
||||||
|
}),
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_alerts_dropped_total",
|
||||||
|
Help: "Alerts dropped, the oldest first, from a full queue, and " +
|
||||||
|
"alerts given up as the destination refused them.",
|
||||||
|
ConstLabels: destination,
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(queue.Counts(name).Dropped)
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// ServeHTTP answers with the metrics in the Prometheus text format.
|
// ServeHTTP answers with the metrics in the Prometheus text format.
|
||||||
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
m.handler.ServeHTTP(w, r)
|
m.handler.ServeHTTP(w, r)
|
||||||
@@ -255,7 +327,15 @@ func (m *Metrics) RequestEnded(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if line.LimitHit != "" {
|
if line.LimitHit != "" {
|
||||||
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
|
// The log line names a byte limit's window with _bytes after it.
|
||||||
|
window, isBytes := strings.CutSuffix(line.LimitHit, "_bytes")
|
||||||
|
|
||||||
|
kind := ratelimit.KindRequests
|
||||||
|
if isBytes {
|
||||||
|
kind = ratelimit.KindBytes
|
||||||
|
}
|
||||||
|
|
||||||
|
m.rateLimitHits.WithLabelValues(window, kind).Inc()
|
||||||
}
|
}
|
||||||
|
|
||||||
if limit != "" {
|
if limit != "" {
|
||||||
@@ -267,7 +347,11 @@ func (m *Metrics) RequestEnded(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if line.Country != "" {
|
if line.Country != "" {
|
||||||
m.countries.add(line)
|
m.countries.add(line.Country, line)
|
||||||
|
}
|
||||||
|
|
||||||
|
if line.ASN != "" {
|
||||||
|
m.asns.add(line.ASN, line)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+50
-7
@@ -1,6 +1,7 @@
|
|||||||
package proxy
|
package proxy
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"crypto/subtle"
|
"crypto/subtle"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -8,6 +9,7 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -27,8 +29,13 @@ const banBodyMaxBytes = 4 << 10
|
|||||||
const permanent = "permanent"
|
const permanent = "permanent"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
errNotBanToAdd = errors.New(
|
||||||
|
"the body is not a JSON object of netblock, duration and reason")
|
||||||
errNotNetblock = errors.New(
|
errNotNetblock = errors.New(
|
||||||
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
|
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
|
||||||
|
errMappedNetblock = errors.New(
|
||||||
|
"is IPv4-mapped: give the IPv4 netblock, such as 203.0.113.0/24")
|
||||||
|
errZone = errors.New("has a zone, which a netblock cannot have")
|
||||||
errNotDuration = errors.New(
|
errNotDuration = errors.New(
|
||||||
"is not a duration above zero, such as 1h or 7d, or permanent")
|
"is not a duration above zero, such as 1h or 7d, or permanent")
|
||||||
errNotAddress = errors.New("is not an address, such as 203.0.113.9")
|
errNotAddress = errors.New("is not an address, such as 203.0.113.9")
|
||||||
@@ -111,6 +118,10 @@ type banToAdd struct {
|
|||||||
// an admin, from now for the duration the body gives, with its reason,
|
// an admin, from now for the duration the body gives, with its reason,
|
||||||
// and answers with that ban.
|
// and answers with that ban.
|
||||||
func (rq *request) addBan() {
|
func (rq *request) addBan() {
|
||||||
|
// The body must arrive within SWWAF_CLIENT_REQUEST_TIMEOUT, as any
|
||||||
|
// other request's must.
|
||||||
|
rq.stopReadingBody(rq.clientRequestDeadline())
|
||||||
|
|
||||||
toAdd, err := rq.readBanToAdd()
|
toAdd, err := rq.readBanToAdd()
|
||||||
if refused := rq.refused.Load(); refused != nil {
|
if refused := rq.refused.Load(); refused != nil {
|
||||||
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
|
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
|
||||||
@@ -118,6 +129,16 @@ func (rq *request) addBan() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if errors.Is(err, os.ErrDeadlineExceeded) {
|
||||||
|
rq.answer(refusal{
|
||||||
|
status: http.StatusRequestTimeout,
|
||||||
|
action: requestlog.ActionTimedOut,
|
||||||
|
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
|
||||||
|
})
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
netblock netip.Prefix
|
netblock netip.Prefix
|
||||||
expires time.Time
|
expires time.Time
|
||||||
@@ -142,23 +163,33 @@ func (rq *request) addBan() {
|
|||||||
rq.answerBans([]bans.Ban{ban})
|
rq.answerBans([]bans.Ban{ban})
|
||||||
}
|
}
|
||||||
|
|
||||||
// readBanToAdd reads the body of POST BansPath, at most banBodyMaxBytes
|
// readBanToAdd reads the body of POST BansPath: a JSON object with
|
||||||
// of it.
|
// nothing but whitespace after it, in at most banBodyMaxBytes.
|
||||||
func (rq *request) readBanToAdd() (banToAdd, error) {
|
func (rq *request) readBanToAdd() (banToAdd, error) {
|
||||||
var body io.ReadCloser = http.NoBody
|
var body io.ReadCloser = http.NoBody
|
||||||
if rq.body != nil {
|
if rq.body != nil {
|
||||||
body = rq.body
|
body = rq.body
|
||||||
}
|
}
|
||||||
|
|
||||||
|
data, err := io.ReadAll(http.MaxBytesReader(nil, body, banBodyMaxBytes))
|
||||||
|
if err != nil {
|
||||||
|
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
|
||||||
|
}
|
||||||
|
|
||||||
var toAdd banToAdd
|
var toAdd banToAdd
|
||||||
|
|
||||||
decoder := json.NewDecoder(http.MaxBytesReader(nil, body, banBodyMaxBytes))
|
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||||
decoder.DisallowUnknownFields()
|
decoder.DisallowUnknownFields()
|
||||||
|
|
||||||
err := decoder.Decode(&toAdd)
|
err = decoder.Decode(&toAdd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return banToAdd{}, fmt.Errorf(
|
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
|
||||||
"the body is not a JSON object of netblock, duration and reason: %w", err)
|
}
|
||||||
|
|
||||||
|
// Token returns io.EOF only when nothing but whitespace is left.
|
||||||
|
_, err = decoder.Token()
|
||||||
|
if !errors.Is(err, io.EOF) {
|
||||||
|
return banToAdd{}, fmt.Errorf("%w: more follows the object", errNotBanToAdd)
|
||||||
}
|
}
|
||||||
|
|
||||||
return toAdd, nil
|
return toAdd, nil
|
||||||
@@ -166,18 +197,30 @@ func (rq *request) readBanToAdd() (banToAdd, error) {
|
|||||||
|
|
||||||
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a
|
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a
|
||||||
// client's address, which stands for the netblock a ban on that client
|
// client's address, which stands for the netblock a ban on that client
|
||||||
// covers.
|
// covers. An IPv4-mapped netblock, such as ::ffff:203.0.113.0/120, is
|
||||||
|
// refused, since a client's address is looked up as IPv4 and a ban on it
|
||||||
|
// would refuse nothing, and so is a value with a zone.
|
||||||
func (h *handler) banNetblock(value string) (netip.Prefix, error) {
|
func (h *handler) banNetblock(value string) (netip.Prefix, error) {
|
||||||
netblock, err := netip.ParsePrefix(value)
|
netblock, err := netip.ParsePrefix(value)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
if netblock.Addr().Is4In6() {
|
||||||
|
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errMappedNetblock)
|
||||||
|
}
|
||||||
|
|
||||||
return netblock, nil
|
return netblock, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ParsePrefix refuses a zone, but ParseAddr reads the /48 of
|
||||||
|
// 2001:db8::1%x/48 as part of the zone.
|
||||||
addr, err := netip.ParseAddr(value)
|
addr, err := netip.ParseAddr(value)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
|
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if addr.Zone() != "" {
|
||||||
|
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errZone)
|
||||||
|
}
|
||||||
|
|
||||||
return h.netblock(addr), nil
|
return h.netblock(addr), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -173,7 +173,9 @@ func TestBanToAddGivesItsNetblockAndDuration(t *testing.T) {
|
|||||||
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
|
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
|
||||||
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
|
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
|
||||||
} {
|
} {
|
||||||
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}`
|
// Whitespace may follow the object.
|
||||||
|
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` +
|
||||||
|
"\r\n"
|
||||||
want := state.BanEntry{
|
want := state.BanEntry{
|
||||||
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
|
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
|
||||||
}
|
}
|
||||||
@@ -203,6 +205,29 @@ func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
|
|||||||
`{"netblock": "203.0.113", "duration": "1h"}`,
|
`{"netblock": "203.0.113", "duration": "1h"}`,
|
||||||
`netblock "203.0.113" is not an address or a netblock`,
|
`netblock "203.0.113" is not an address or a netblock`,
|
||||||
},
|
},
|
||||||
|
// A client's address is looked up as IPv4, so a ban on an
|
||||||
|
// IPv4-mapped netblock would refuse nothing.
|
||||||
|
{
|
||||||
|
`{"netblock": "::ffff:203.0.113.0/120", "duration": "1h"}`,
|
||||||
|
`netblock "::ffff:203.0.113.0/120" is IPv4-mapped`,
|
||||||
|
},
|
||||||
|
// Read as an address, its zone would be "x/48", and its ban on the
|
||||||
|
// /64 around it.
|
||||||
|
{
|
||||||
|
`{"netblock": "2001:db8::1%x/48", "duration": "1h"}`,
|
||||||
|
`netblock "2001:db8::1%x/48" has a zone`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
`{"netblock": "fe80::1%eth0", "duration": "1h"}`,
|
||||||
|
`netblock "fe80::1%eth0" has a zone`,
|
||||||
|
},
|
||||||
|
// Anything but whitespace after the object.
|
||||||
|
{
|
||||||
|
`{"netblock": "203.0.113.9", "duration": "1h"}` +
|
||||||
|
`{"netblock": "198.51.100.0/24", "duration": "1h"}`,
|
||||||
|
"more follows the object",
|
||||||
|
},
|
||||||
|
{`{"netblock": "203.0.113.9", "duration": "1h"} x`, "more follows the object"},
|
||||||
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
|
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
|
||||||
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
|
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
|
||||||
{
|
{
|
||||||
@@ -218,12 +243,16 @@ func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
|
|||||||
`duration "forever" is not a duration above zero, such as 1h or 7d, ` +
|
`duration "forever" is not a duration above zero, such as 1h or 7d, ` +
|
||||||
`or permanent`,
|
`or permanent`,
|
||||||
},
|
},
|
||||||
// Over the 4 KiB read of a body.
|
// Over the 4 KiB read of a body, even when the object comes first.
|
||||||
{
|
{
|
||||||
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
|
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
|
||||||
strings.Repeat("x", 4<<10) + `"}`,
|
strings.Repeat("x", 4<<10) + `"}`,
|
||||||
"request body too large",
|
"request body too large",
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
`{"netblock": "203.0.113.9", "duration": "1h"}` + strings.Repeat(" ", 4<<10),
|
||||||
|
"request body too large",
|
||||||
|
},
|
||||||
} {
|
} {
|
||||||
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
|
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
|
||||||
if !strings.Contains(string(got.body), tc.want) {
|
if !strings.Contains(string(got.body), tc.want) {
|
||||||
@@ -257,6 +286,36 @@ func TestBanToAddOverTheRequestSizeLimitIsRefused(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBanToAddSlowerThanTheClientRequestTimeoutIsRefused(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, server := startWithClock(t, "", map[string]string{
|
||||||
|
adminToken: adminSecret,
|
||||||
|
metricsToken: token,
|
||||||
|
clientRequestTimeout: shortTimeoutSetting,
|
||||||
|
})
|
||||||
|
|
||||||
|
// The chunk announces 256 bytes and the rest of it never comes, so only
|
||||||
|
// the timeout ends the wait. A hold-up of the test process can only
|
||||||
|
// make the answer later, so the time is checked only for not being
|
||||||
|
// shorter than the timeout.
|
||||||
|
start := time.Now()
|
||||||
|
|
||||||
|
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
|
||||||
|
http.MethodPost, proxy.BansPath, "100\r\n"+`{"netblock": "203.0.113.9", `,
|
||||||
|
http.StatusRequestTimeout, requestlog.ActionTimedOut)
|
||||||
|
|
||||||
|
if took := time.Since(start); took < shortTimeout {
|
||||||
|
t.Errorf("answered after %s, before the timeout of %s ran out", took, shortTimeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantLimitHits(t, s.addr, clientRequestTimeout, 1)
|
||||||
|
|
||||||
|
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
||||||
|
t.Errorf("the ledger holds %+v, want no ban", held)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
|
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,273 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"maps"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
|
||||||
|
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
|
||||||
|
// alertInstance is the instance every alert of these tests gives.
|
||||||
|
alertInstance = "fsn1app1/gitea"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
banScopeV4Prefix: "24",
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
||||||
|
ban := server.Ledger.Bans(netblock)[0]
|
||||||
|
|
||||||
|
// A request refused under the ban raises no other alert.
|
||||||
|
clk.advance(time.Minute)
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{
|
||||||
|
Netblock: netblock, Cause: bans.CauseLimit,
|
||||||
|
Reason: "requests per minute over the limit of 1", Notes: ban.Notes,
|
||||||
|
}, requestlog.FormatTime(start.Add(time.Hour))))
|
||||||
|
|
||||||
|
if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" {
|
||||||
|
t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
netblock := netip.MustParsePrefix(client + "/32")
|
||||||
|
other := netip.MustParsePrefix(otherClient + "/32")
|
||||||
|
|
||||||
|
// The probe bans the client for seven days, and its next request makes
|
||||||
|
// the ban permanent. The request after that changes nothing.
|
||||||
|
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
attackBan := server.Ledger.Bans(netblock)[0]
|
||||||
|
|
||||||
|
clk.advance(time.Minute)
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
permanentBan := server.Ledger.Bans(netblock)[0]
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
// Another client's probe after its first ban has run out without a
|
||||||
|
// request makes a permanent ban at once.
|
||||||
|
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
clk.advance(7 * 24 * time.Hour)
|
||||||
|
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
otherBans := server.Ledger.Bans(other)
|
||||||
|
|
||||||
|
wantAlerts(t, queue,
|
||||||
|
attackAlert(alerts.EventBan, start, client, attackBan,
|
||||||
|
requestlog.FormatTime(start.Add(7*24*time.Hour))),
|
||||||
|
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute), client,
|
||||||
|
permanentBan, "permanent"),
|
||||||
|
attackAlert(alerts.EventBan, start.Add(time.Minute), otherClient, otherBans[0],
|
||||||
|
requestlog.FormatTime(start.Add(time.Minute+7*24*time.Hour))),
|
||||||
|
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute+7*24*time.Hour),
|
||||||
|
otherClient, otherBans[1], "permanent"),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
rateLimitPerMinute: "2",
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
// A ban for a clear sign of attack, which a request under it would make
|
||||||
|
// permanent.
|
||||||
|
group := netip.MustParsePrefix(ipv6Group)
|
||||||
|
attackBan, _ := server.Ledger.BanForAttack(group, start, bans.Notes{RuleID: "probe"})
|
||||||
|
|
||||||
|
// The third request breaks the limit, and so does the fourth, within the
|
||||||
|
// cooldown, which raises nothing. The probe is a clear sign of attack.
|
||||||
|
for range 4 {
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
|
||||||
|
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
// No ban is made, and none made permanent.
|
||||||
|
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan ||
|
||||||
|
line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
|
||||||
|
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+
|
||||||
|
"for the attack alone, as it was", held, line.BanExpires)
|
||||||
|
}
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 3 || queue.Suppressed() != 0 {
|
||||||
|
t.Fatalf("%d alerts wait and %d are held back, want 3 and 0: %+v",
|
||||||
|
len(waiting), queue.Suppressed(), waiting)
|
||||||
|
}
|
||||||
|
|
||||||
|
limitNotes, _ := waiting[0].Detail["notes"].(bans.Notes)
|
||||||
|
attackNotes, _ := waiting[1].Detail["notes"].(bans.Notes)
|
||||||
|
|
||||||
|
if limitNotes.Limit != 2 || limitNotes.Request.Path != "/" ||
|
||||||
|
attackNotes.Request.Path != "/.env" {
|
||||||
|
t.Errorf("the notes are %+v and %+v, want those of the broken limit and "+
|
||||||
|
"of the probe", limitNotes, attackNotes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Each alert is the one enforce mode would have raised, with mode
|
||||||
|
// observe in its detail.
|
||||||
|
want := []alerts.Alert{
|
||||||
|
banAlert(alerts.EventBan, start, client, bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
|
||||||
|
Reason: "requests per minute over the limit of 2", Notes: limitNotes,
|
||||||
|
}, requestlog.FormatTime(start.Add(time.Hour))),
|
||||||
|
attackAlert(alerts.EventBan, start, otherClient, bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix(otherClient + "/32"), Notes: attackNotes,
|
||||||
|
}, requestlog.FormatTime(start.Add(7*24*time.Hour))),
|
||||||
|
attackAlert(alerts.EventPermanentBan, start, ipv6Client, attackBan, permanent),
|
||||||
|
}
|
||||||
|
for _, alert := range want {
|
||||||
|
alert.Detail["mode"] = observe
|
||||||
|
}
|
||||||
|
|
||||||
|
wantAlerts(t, queue, want...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeWorksOutABanOnlyWhenItsAlertWouldBeSent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _, queue := startWithAlerts(t, map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
rateLimitPerMinute: "2",
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
alertMaxPerHour: "2",
|
||||||
|
})
|
||||||
|
|
||||||
|
// The client's third request breaks the limit, and raises the first
|
||||||
|
// alert of the hour. Its fourth is within the cooldown.
|
||||||
|
for range 4 {
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The other client's first probe raises the second. Its second probe is
|
||||||
|
// within the cooldown.
|
||||||
|
for range 2 {
|
||||||
|
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The IPv6 client's third request breaks the limit past the two alerts
|
||||||
|
// an hour.
|
||||||
|
for range 3 {
|
||||||
|
s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Had the ban been worked out for any of the requests within the
|
||||||
|
// cooldown or past the two an hour, its alert would have been raised,
|
||||||
|
// held back and counted.
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 2 || queue.Suppressed() != 0 {
|
||||||
|
t.Errorf("%d alerts wait and %d are held back, want 2 and 0: %+v",
|
||||||
|
len(waiting), queue.Suppressed(), waiting)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// startWithAlerts is startWithClock with alerts to a webhook, which is
|
||||||
|
// never sent them, and returns the queue they wait in as well.
|
||||||
|
func startWithAlerts(
|
||||||
|
t *testing.T, env map[string]string,
|
||||||
|
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
return startAppWithAlerts(t, func(http.ResponseWriter, *http.Request) {}, env)
|
||||||
|
}
|
||||||
|
|
||||||
|
// startAppWithAlerts is startWithAlerts in front of the app handler.
|
||||||
|
func startAppWithAlerts(
|
||||||
|
t *testing.T, handler http.HandlerFunc, env map[string]string,
|
||||||
|
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
app := startApp(t, handler)
|
||||||
|
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
|
||||||
|
settings := map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
alertWebhookURL: "https://alerts.example/smallwebwaf",
|
||||||
|
instanceName: alertInstance,
|
||||||
|
}
|
||||||
|
maps.Copy(settings, env)
|
||||||
|
|
||||||
|
addr, out, server, queue := startProxyWithAlerts(t, app.URL, "", clk.Now, settings)
|
||||||
|
|
||||||
|
return &sender{t: t, addr: addr, out: out}, clk, server, queue
|
||||||
|
}
|
||||||
|
|
||||||
|
// banAlert returns the alert for event, raised by a request from client at
|
||||||
|
// the time raised, for ban, with its netblock, cause, reason and notes,
|
||||||
|
// which ends at expires, as the log line gives it.
|
||||||
|
func banAlert(
|
||||||
|
event string, raised time.Time, client string, ban bans.Ban, expires string,
|
||||||
|
) alerts.Alert {
|
||||||
|
return alerts.Alert{
|
||||||
|
Instance: alertInstance,
|
||||||
|
Time: raised,
|
||||||
|
Event: event,
|
||||||
|
Client: netip.MustParseAddr(client),
|
||||||
|
Netblock: ban.Netblock,
|
||||||
|
Reason: ban.Reason,
|
||||||
|
Detail: map[string]any{
|
||||||
|
"cause": ban.Cause, "ban_expires": expires, "notes": ban.Notes,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// attackAlert is banAlert for a ban for the probe rule of testRules, with
|
||||||
|
// the netblock and the notes of ban.
|
||||||
|
func attackAlert(
|
||||||
|
event string, raised time.Time, client string, ban bans.Ban, expires string,
|
||||||
|
) alerts.Alert {
|
||||||
|
return banAlert(event, raised, client, bans.Ban{
|
||||||
|
Netblock: ban.Netblock, Cause: bans.CauseAttack, Reason: "matched the rule probe",
|
||||||
|
Notes: ban.Notes,
|
||||||
|
}, expires)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantAlerts checks the alerts waiting in queue, in order.
|
||||||
|
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(got) != len(want) {
|
||||||
|
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := range want {
|
||||||
|
if !reflect.DeepEqual(got[i], want[i]) {
|
||||||
|
t.Errorf("alert %d is\n%+v\nwant\n%+v", i, got[i], want[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+191
-28
@@ -4,7 +4,9 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
)
|
)
|
||||||
@@ -16,81 +18,242 @@ func (rq *request) banResponse(action string) *refusal {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// banned reports whether a ban on a netblock the client is in covers the
|
// banned reports whether a ban on a netblock the client is in covers the
|
||||||
// request at now, and notes for the log line when that ban ends.
|
// request at now, and notes for the log line when that ban ends. A
|
||||||
|
// request that makes the ban permanent, or in observe mode would have,
|
||||||
|
// raises the alert for it.
|
||||||
func (rq *request) banned(now time.Time) bool {
|
func (rq *request) banned(now time.Time) bool {
|
||||||
check := rq.h.ledger.Check
|
check := rq.h.ledger.Check
|
||||||
if rq.h.config.Observe {
|
if rq.h.config.Observe {
|
||||||
check = rq.h.ledger.Find // in observe mode the ban refuses nothing
|
check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is
|
||||||
}
|
}
|
||||||
|
|
||||||
ban, banned := check(rq.client, now)
|
ban, banned, madePermanent := check(rq.client, now)
|
||||||
if banned {
|
if banned {
|
||||||
rq.line.BanExpires = banExpires(ban)
|
rq.line.BanExpires = banExpires(ban)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if madePermanent {
|
||||||
|
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
|
||||||
|
rq.alertBan(ban)
|
||||||
|
}
|
||||||
|
|
||||||
return banned
|
return banned
|
||||||
}
|
}
|
||||||
|
|
||||||
// limitBroken counts the request for the rate limits at now, notes the
|
// limitBroken counts the request for the rate limits at now, notes the
|
||||||
// client's counts for the log line, and reports whether the request takes
|
// client's counts for the log line, and reports whether the request takes
|
||||||
// the client over a limit. In enforce mode such a request bans the
|
// the client over a rate limit, as its limit percentage lowers it, which
|
||||||
// client's netblock, and sets the client's counters back to zero; in
|
// breaks it.
|
||||||
// observe mode it does neither.
|
|
||||||
func (rq *request) limitBroken(now time.Time) bool {
|
func (rq *request) limitBroken(now time.Time) bool {
|
||||||
group := clientGroup(rq.client)
|
counts, hit, over := rq.h.limiter.Count(clientGroup(rq.client), now,
|
||||||
|
rq.limitPercent.percent)
|
||||||
counts, hit, over := rq.h.limiter.Count(group, now)
|
|
||||||
rq.line.Counts = counts
|
rq.line.Counts = counts
|
||||||
|
|
||||||
if !over {
|
if over {
|
||||||
return false
|
rq.banForLimit(now, hit, rq.h.config.BanResponse)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return over
|
||||||
|
}
|
||||||
|
|
||||||
|
// countBytes counts the request's bytes for the byte limits, once its
|
||||||
|
// response has ended, and notes the client's byte totals for the log line;
|
||||||
|
// its requests stay there as the rate limits counted them. The bytes are
|
||||||
|
// the response's body bytes, the request's, or both, as SWWAF_BYTES_COUNT
|
||||||
|
// says; for an upgraded connection, such as a WebSocket, which has closed
|
||||||
|
// by then, what it carried from the app counts with the response's and
|
||||||
|
// what it carried from the client with the request's. Only a request
|
||||||
|
// passed to the app has them counted, and only one the rate limits
|
||||||
|
// counted; in observe mode, not one that enforce mode would have refused.
|
||||||
|
// Bytes that take the client over a byte limit, as its limit percentage
|
||||||
|
// for the byte limits lowers it, break it; the response was passed on
|
||||||
|
// whole.
|
||||||
|
func (rq *request) countBytes() {
|
||||||
|
if !rq.counted || rq.line.WouldAction != "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
response, request := rq.out.bytes, rq.requestBytes()
|
||||||
|
if rq.upgraded != nil {
|
||||||
|
response += rq.upgraded.fromApp.Load()
|
||||||
|
request += rq.upgraded.toApp.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
var bytes int64
|
||||||
|
|
||||||
|
switch rq.h.config.BytesCount {
|
||||||
|
case "response":
|
||||||
|
bytes = response
|
||||||
|
case "request":
|
||||||
|
bytes = request
|
||||||
|
default: // both
|
||||||
|
bytes = response + request
|
||||||
|
}
|
||||||
|
|
||||||
|
now := rq.h.now()
|
||||||
|
|
||||||
|
counts, hit, over := rq.h.limiter.CountBytes(clientGroup(rq.client), now, bytes,
|
||||||
|
rq.bytesPercent.percent)
|
||||||
|
rq.line.Counts.MinuteBytes = counts.MinuteBytes
|
||||||
|
rq.line.Counts.HourBytes = counts.HourBytes
|
||||||
|
rq.line.Counts.DayBytes = counts.DayBytes
|
||||||
|
|
||||||
|
if over {
|
||||||
|
rq.banForLimit(now, hit, rq.out.status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// banForLimit bans the client's netblock at now for a broken limit, the
|
||||||
|
// one hit names, and notes the offence for the log line. status is what
|
||||||
|
// the client was sent, or is sent: SWWAF_BAN_RESPONSE for a request over
|
||||||
|
// a rate limit, the app's answer for one whose bytes broke a byte limit.
|
||||||
|
// The ban's notes give the client's limit percentage for that kind of
|
||||||
|
// limit. The ban sets the client's counters back to zero. In observe mode
|
||||||
|
// it makes no ban and sets nothing back, and raises the alert for the ban
|
||||||
|
// it would have made, if that alert would be sent.
|
||||||
|
func (rq *request) banForLimit(now time.Time, hit ratelimit.Hit, status int) {
|
||||||
rq.line.LimitHit = hit.Window
|
rq.line.LimitHit = hit.Window
|
||||||
|
if hit.Kind == ratelimit.KindBytes {
|
||||||
|
rq.line.LimitHit += "_bytes" // as counts names the byte totals
|
||||||
|
}
|
||||||
|
|
||||||
rq.line.Offence = requestlog.OffenceLimit
|
rq.line.Offence = requestlog.OffenceLimit
|
||||||
|
|
||||||
if rq.h.config.Observe {
|
netblock := rq.h.netblock(rq.client)
|
||||||
return true
|
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
netblock := rq.h.netblock(rq.client)
|
notes := bans.Notes{
|
||||||
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
|
ASN: rq.line.ASN,
|
||||||
|
ASName: rq.line.ASName,
|
||||||
Country: rq.line.Country,
|
Country: rq.line.Country,
|
||||||
|
Kind: hit.Kind,
|
||||||
Limit: hit.Limit,
|
Limit: hit.Limit,
|
||||||
Window: hit.Window,
|
Window: hit.Window,
|
||||||
Count: hit.Requests,
|
Count: hit.Count,
|
||||||
Request: rq.noted(now),
|
Request: rq.noted(now, status),
|
||||||
Requests: rq.netblockRequests(netblock),
|
Requests: rq.netblockRequests(netblock),
|
||||||
})
|
}
|
||||||
rq.h.limiter.Reset(group)
|
|
||||||
|
percent := rq.limitPercent
|
||||||
|
if hit.Kind == ratelimit.KindBytes {
|
||||||
|
percent = rq.bytesPercent
|
||||||
|
}
|
||||||
|
|
||||||
|
notes.LimitPercent, notes.LimitPercentSetting = percent.logged()
|
||||||
|
|
||||||
|
if rq.h.config.Observe {
|
||||||
|
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
|
||||||
|
if wouldBan {
|
||||||
|
rq.alertBan(ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
|
||||||
|
rq.h.limiter.Reset(clientGroup(rq.client))
|
||||||
rq.line.BanExpires = banExpires(ban)
|
rq.line.BanExpires = banExpires(ban)
|
||||||
|
|
||||||
return true
|
if made {
|
||||||
|
rq.alertBan(ban)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// banForAttack bans the client's netblock at now for a clear sign of
|
// banForAttack bans the client's netblock at now for a clear sign of
|
||||||
// attack, the match of rule, a ban rule.
|
// attack, the match of rule, a ban rule. In observe mode it makes no ban,
|
||||||
|
// and raises the alert for the ban it would have made, if that alert
|
||||||
|
// would be sent.
|
||||||
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
|
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
|
||||||
netblock := rq.h.netblock(rq.client)
|
netblock := rq.h.netblock(rq.client)
|
||||||
ban := rq.h.ledger.BanForAttack(netblock, now, bans.Notes{
|
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := bans.Notes{
|
||||||
|
ASN: rq.line.ASN,
|
||||||
|
ASName: rq.line.ASName,
|
||||||
Country: rq.line.Country,
|
Country: rq.line.Country,
|
||||||
RuleID: rule.ID,
|
RuleID: rule.ID,
|
||||||
Target: rule.Target,
|
Target: rule.Target,
|
||||||
Request: rq.noted(now),
|
Request: rq.noted(now, rq.h.config.BanResponse),
|
||||||
Requests: rq.netblockRequests(netblock),
|
Requests: rq.netblockRequests(netblock),
|
||||||
})
|
}
|
||||||
|
|
||||||
|
if rq.h.config.Observe {
|
||||||
|
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
|
||||||
|
if wouldBan {
|
||||||
|
rq.alertBan(ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ban, made := rq.h.ledger.BanForAttack(netblock, now, notes)
|
||||||
rq.line.BanExpires = banExpires(ban)
|
rq.line.BanExpires = banExpires(ban)
|
||||||
|
|
||||||
|
if made {
|
||||||
|
rq.alertBan(ban)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, as the
|
// wouldAlertBan reports whether the alert for a ban on netblock for cause
|
||||||
// notes of the ban it makes keep it.
|
// made at now would be sent. In observe mode the ban the request would
|
||||||
func (rq *request) noted(now time.Time) bans.Request {
|
// have made is worked out only then, at most once per
|
||||||
|
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
|
||||||
|
// netblock's requests, which can mean going through every client.
|
||||||
|
func (rq *request) wouldAlertBan(
|
||||||
|
netblock netip.Prefix, now time.Time, cause string,
|
||||||
|
) bool {
|
||||||
|
event := alerts.EventBan
|
||||||
|
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
|
||||||
|
event = alerts.EventPermanentBan
|
||||||
|
}
|
||||||
|
|
||||||
|
return rq.h.alerts.WouldSend(event, netblock)
|
||||||
|
}
|
||||||
|
|
||||||
|
// alertBan raises the alert for ban, which the request made, or made
|
||||||
|
// permanent: permanent_ban for a permanent ban, ban for another. Its
|
||||||
|
// detail gives the ban's cause, when it ends, and its notes, and in
|
||||||
|
// observe mode, where ban is the ban that would have been made, or made
|
||||||
|
// permanent, mode, observe.
|
||||||
|
func (rq *request) alertBan(ban bans.Ban) {
|
||||||
|
event := alerts.EventBan
|
||||||
|
if ban.Permanent() {
|
||||||
|
event = alerts.EventPermanentBan
|
||||||
|
}
|
||||||
|
|
||||||
|
detail := map[string]any{
|
||||||
|
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
|
||||||
|
}
|
||||||
|
if rq.h.config.Observe {
|
||||||
|
detail["mode"] = "observe"
|
||||||
|
}
|
||||||
|
|
||||||
|
rq.h.alerts.Raise(alerts.Alert{
|
||||||
|
Event: event,
|
||||||
|
Client: rq.client,
|
||||||
|
Netblock: ban.Netblock,
|
||||||
|
ASN: ban.Notes.ASN,
|
||||||
|
ASName: ban.Notes.ASName,
|
||||||
|
Country: ban.Notes.Country,
|
||||||
|
Reason: ban.Reason,
|
||||||
|
Detail: detail,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// noted is the request, at now, with status, what the client was sent, or
|
||||||
|
// in observe mode would have been, as the notes of the ban it makes keep
|
||||||
|
// it.
|
||||||
|
func (rq *request) noted(now time.Time, status int) bans.Request {
|
||||||
return bans.Request{
|
return bans.Request{
|
||||||
Time: now,
|
Time: now,
|
||||||
Method: rq.in.Method,
|
Method: rq.in.Method,
|
||||||
Host: rq.in.Host,
|
Host: rq.in.Host,
|
||||||
Path: rq.in.URL.RequestURI(),
|
Path: rq.in.URL.RequestURI(),
|
||||||
Status: rq.h.config.BanResponse,
|
Status: status,
|
||||||
UserAgent: rq.in.UserAgent(),
|
UserAgent: rq.in.UserAgent(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -281,7 +281,10 @@ func TestBanNotes(t *testing.T) {
|
|||||||
Cause: bans.CauseLimit,
|
Cause: bans.CauseLimit,
|
||||||
Reason: "requests per minute over the limit of 1",
|
Reason: "requests per minute over the limit of 1",
|
||||||
Notes: bans.Notes{
|
Notes: bans.Notes{
|
||||||
|
ASN: asnDE,
|
||||||
|
ASName: asNameDE,
|
||||||
Country: "DE",
|
Country: "DE",
|
||||||
|
Kind: "requests",
|
||||||
Limit: 1,
|
Limit: 1,
|
||||||
Window: minute,
|
Window: minute,
|
||||||
Count: 2,
|
Count: 2,
|
||||||
@@ -362,8 +365,9 @@ func (c *clock) advance(d time.Duration) {
|
|||||||
|
|
||||||
// startWithClock starts smallwebwaf in front of an app that answers 200,
|
// startWithClock starts smallwebwaf in front of an app that answers 200,
|
||||||
// with the settings in env on top of trusting localhost's
|
// with the settings in env on top of trusting localhost's
|
||||||
// X-Forwarded-For, clients' countries looked up at geojsURL, and a clock
|
// X-Forwarded-For, clients' AS numbers and countries looked up at
|
||||||
// set to midnight, the start of a bucket in every window.
|
// geojsURL, and a clock set to midnight, the start of a bucket in every
|
||||||
|
// window.
|
||||||
func startWithClock(
|
func startWithClock(
|
||||||
t *testing.T, geojsURL string, env map[string]string,
|
t *testing.T, geojsURL string, env map[string]string,
|
||||||
) (*sender, *clock, *proxy.Server) {
|
) (*sender, *clock, *proxy.Server) {
|
||||||
|
|||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// whole is the percentage of each limit a client gets when no biased
|
||||||
|
// threshold lowers its limits.
|
||||||
|
const whole = 100
|
||||||
|
|
||||||
|
// percentage is a client's limit percentage for the rate limits or for
|
||||||
|
// the byte limits, as the biased thresholds give it, and the setting that
|
||||||
|
// gave it: "" with whole when none lowers that kind of limit.
|
||||||
|
type percentage struct {
|
||||||
|
percent int64
|
||||||
|
setting string
|
||||||
|
}
|
||||||
|
|
||||||
|
// biasedThresholdsSet reports whether a biased threshold can lower a
|
||||||
|
// client's limits: one of its lists is not empty, or
|
||||||
|
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100. The client's lookup is then
|
||||||
|
// needed before its request goes on.
|
||||||
|
func biasedThresholdsSet(cfg *config.Config) bool {
|
||||||
|
return len(cfg.ASNLimitPercent) > 0 || len(cfg.CountryLimitPercent) > 0 ||
|
||||||
|
len(cfg.ASNBytesPercent) > 0 || len(cfg.CountryBytesPercent) > 0 ||
|
||||||
|
cfg.UnknownLimitPercent < whole
|
||||||
|
}
|
||||||
|
|
||||||
|
// limitPercentages returns a client's limit percentages, for the rate
|
||||||
|
// limits and for the byte limits, by its AS number and country as looked
|
||||||
|
// up, each "" when unknown. Each is the lowest of those the settings give
|
||||||
|
// it, the first of them in the order below when several are lowest: the
|
||||||
|
// percentage SWWAF_ASN_LIMIT_PERCENT gives its AS number, the one
|
||||||
|
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, and, for a client
|
||||||
|
// without a country, SWWAF_UNKNOWN_LIMIT_PERCENT. For the byte limits,
|
||||||
|
// SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT take the place
|
||||||
|
// of the first two for an AS number or a country they list.
|
||||||
|
func limitPercentages(
|
||||||
|
cfg *config.Config, asn, country string,
|
||||||
|
) (percentage, percentage) {
|
||||||
|
unknown := percentage{percent: whole}
|
||||||
|
if country == "" {
|
||||||
|
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
|
||||||
|
}
|
||||||
|
|
||||||
|
asnRequests := given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT")
|
||||||
|
countryRequests := given(cfg.CountryLimitPercent, country,
|
||||||
|
"SWWAF_COUNTRY_LIMIT_PERCENT")
|
||||||
|
|
||||||
|
asnBytes, countryBytes := asnRequests, countryRequests
|
||||||
|
if _, listed := cfg.ASNBytesPercent[asn]; listed {
|
||||||
|
asnBytes = given(cfg.ASNBytesPercent, asn, "SWWAF_ASN_BYTES_PERCENT")
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, listed := cfg.CountryBytesPercent[country]; listed {
|
||||||
|
countryBytes = given(cfg.CountryBytesPercent, country, "SWWAF_COUNTRY_BYTES_PERCENT")
|
||||||
|
}
|
||||||
|
|
||||||
|
return lowest(asnRequests, countryRequests, unknown),
|
||||||
|
lowest(asnBytes, countryBytes, unknown)
|
||||||
|
}
|
||||||
|
|
||||||
|
// given returns the percentage percents, the setting named setting, gives
|
||||||
|
// code, an AS number or a country, or whole when it does not list code.
|
||||||
|
func given(percents map[string]int64, code, setting string) percentage {
|
||||||
|
percent, listed := percents[code]
|
||||||
|
if !listed {
|
||||||
|
return percentage{percent: whole}
|
||||||
|
}
|
||||||
|
|
||||||
|
return percentage{percent, setting}
|
||||||
|
}
|
||||||
|
|
||||||
|
// lowest returns the lowest of percentages below whole, the first of them
|
||||||
|
// when several are lowest, or whole when none is below it.
|
||||||
|
func lowest(percentages ...percentage) percentage {
|
||||||
|
low := percentage{percent: whole}
|
||||||
|
|
||||||
|
for _, p := range percentages {
|
||||||
|
if p.percent < low.percent {
|
||||||
|
low = p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return low
|
||||||
|
}
|
||||||
|
|
||||||
|
// logged returns p as the log line and the notes of a ban give it: its
|
||||||
|
// percent and setting, or nil and "" for whole, which they leave out.
|
||||||
|
func (p percentage) logged() (*int64, string) {
|
||||||
|
if p.percent == whole {
|
||||||
|
return nil, ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return &p.percent, p.setting
|
||||||
|
}
|
||||||
@@ -0,0 +1,494 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"maps"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/netip"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The biased thresholds.
|
||||||
|
const (
|
||||||
|
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
|
||||||
|
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
|
||||||
|
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
|
||||||
|
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
|
||||||
|
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// asnDEHalf and countryDEHalf give fromDE's AS number and its country
|
||||||
|
// half of every limit, and asnDEQuarter gives its AS number a quarter.
|
||||||
|
asnDEHalf = asnDE + ":50"
|
||||||
|
asnDEQuarter = asnDE + ":25"
|
||||||
|
countryDEHalf = "de:50"
|
||||||
|
// noCountry is in an AS of its own, AS64500, and in no country.
|
||||||
|
noCountry = "192.0.2.80"
|
||||||
|
// fourAMinute is the rate limit these tests set: half of it is 2
|
||||||
|
// requests a minute, a quarter of it 1.
|
||||||
|
fourAMinute = "4"
|
||||||
|
// twoUploads is the byte limit these tests set: 199 bytes, which an
|
||||||
|
// upload, a request with a body and its answer, 100 bytes, is within,
|
||||||
|
// and half of which, 99 bytes, it is over.
|
||||||
|
twoUploads = "199"
|
||||||
|
// none is how percentText gives a percentage left out.
|
||||||
|
none = "none"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEachBiasedThresholdLowersTheRateLimits(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting, value, from string
|
||||||
|
}{
|
||||||
|
{asnLimitPercent, asnDEHalf, fromDE},
|
||||||
|
{countryLimitPercent, countryDEHalf, fromDE},
|
||||||
|
{unknownLimitPercent, "50", unplaced},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithLookups(t, map[string]string{
|
||||||
|
rateLimitPerMinute: fourAMinute, tc.setting: tc.value,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Half of 4 requests a minute: the third breaks the limit.
|
||||||
|
for _, sent := range []struct {
|
||||||
|
status int
|
||||||
|
action string
|
||||||
|
}{
|
||||||
|
{http.StatusOK, requestlog.ActionForward},
|
||||||
|
{http.StatusOK, requestlog.ActionForward},
|
||||||
|
{http.StatusForbidden, requestlog.ActionRateLimited},
|
||||||
|
} {
|
||||||
|
line := s.get(tc.from, sent.status, sent.action)
|
||||||
|
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||||
|
"50 from "+tc.setting)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fromKP, which no setting lists, has the whole limit.
|
||||||
|
for range 3 {
|
||||||
|
line := s.get(fromKP, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||||
|
none)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEachBiasedThresholdLowersTheByteLimits(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// The AS numbers and countries are given in either case.
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting, value, from string
|
||||||
|
}{
|
||||||
|
{asnLimitPercent, asnDEHalf, fromDE},
|
||||||
|
{countryLimitPercent, "DE:50", fromDE},
|
||||||
|
{unknownLimitPercent, "50", unplaced},
|
||||||
|
{asnBytesPercent, "as64496:50", fromDE},
|
||||||
|
{countryBytesPercent, countryDEHalf, fromDE},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithLookups(t, map[string]string{
|
||||||
|
bytesLimitPerMinute: twoUploads, tc.setting: tc.value,
|
||||||
|
})
|
||||||
|
|
||||||
|
// The upload's 100 bytes are over half of 199, 99.
|
||||||
|
line := s.uploadFrom(tc.from)
|
||||||
|
if line.LimitHit != minuteBytes {
|
||||||
|
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
|
||||||
|
"50 from "+tc.setting)
|
||||||
|
|
||||||
|
// fromKP, which no setting lists, has the whole limit.
|
||||||
|
line = s.uploadFrom(fromKP)
|
||||||
|
if line.LimitHit != "" {
|
||||||
|
t.Errorf("log line for %s has limit_hit %q, want none", fromKP, line.LimitHit)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting, none)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBytesPercentSettingsTakeThePlaceOfTheOthersForByteLimits(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
env map[string]string
|
||||||
|
// limitPercent and bytesPercent are the log line's, as percentText
|
||||||
|
// gives them, and limitHit is its limit_hit.
|
||||||
|
limitPercent, bytesPercent, limitHit string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"lowering the byte limits alone",
|
||||||
|
map[string]string{asnBytesPercent: asnDEHalf},
|
||||||
|
none, "50 from " + asnBytesPercent, minuteBytes,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"lowering the byte limits alone, by country",
|
||||||
|
map[string]string{countryBytesPercent: countryDEHalf},
|
||||||
|
none, "50 from " + countryBytesPercent, minuteBytes,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"raising the byte limits back",
|
||||||
|
map[string]string{asnLimitPercent: asnDEHalf, asnBytesPercent: asnDE + ":100"},
|
||||||
|
"50 from " + asnLimitPercent, none, "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"raising the byte limits back, by country",
|
||||||
|
map[string]string{countryLimitPercent: countryDEHalf, countryBytesPercent: "de:100"},
|
||||||
|
"50 from " + countryLimitPercent, none, "",
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
env := map[string]string{bytesLimitPerMinute: twoUploads}
|
||||||
|
maps.Copy(env, tc.env)
|
||||||
|
s, _, _ := startWithLookups(t, env)
|
||||||
|
|
||||||
|
// The upload's 100 bytes are over 99, half of 199, and within 199.
|
||||||
|
line := s.uploadFrom(fromDE)
|
||||||
|
if line.LimitHit != tc.limitHit {
|
||||||
|
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||||
|
tc.limitPercent)
|
||||||
|
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
|
||||||
|
tc.bytesPercent)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestZeroPercentIsAZeroAllowance(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithLookups(t, map[string]string{asnLimitPercent: asnDE + ":0"})
|
||||||
|
|
||||||
|
// The first request breaks the limit, and bans the client; the log line
|
||||||
|
// gives the 0.
|
||||||
|
line := s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
if line.fields["limit_percent"] != float64(0) ||
|
||||||
|
line.fields["limit_percent_setting"] != asnLimitPercent {
|
||||||
|
t.Errorf("log line has limit_percent %v from %v, want 0 from %s",
|
||||||
|
line.fields["limit_percent"], line.fields["limit_percent_setting"],
|
||||||
|
asnLimitPercent)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLowestPercentageApplies(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
env map[string]string
|
||||||
|
from string
|
||||||
|
// want is the log line's limit_percent, as percentText gives it.
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"the country's",
|
||||||
|
map[string]string{asnLimitPercent: asnDEHalf, countryLimitPercent: "de:25"},
|
||||||
|
fromDE, "25 from " + countryLimitPercent,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"the AS number's",
|
||||||
|
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: countryDEHalf},
|
||||||
|
fromDE, "25 from " + asnLimitPercent,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"the AS number's, the first of two alike",
|
||||||
|
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: "de:25"},
|
||||||
|
fromDE, "25 from " + asnLimitPercent,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"that for a client without a country",
|
||||||
|
map[string]string{asnLimitPercent: "AS64500:50", unknownLimitPercent: "25"},
|
||||||
|
noCountry, "25 from " + unknownLimitPercent,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// SWWAF_UNKNOWN_LIMIT_PERCENT is left at its default, 100.
|
||||||
|
"the AS number's, for a client without a country",
|
||||||
|
map[string]string{asnLimitPercent: "AS64500:25"},
|
||||||
|
noCountry, "25 from " + asnLimitPercent,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
env := map[string]string{rateLimitPerMinute: fourAMinute}
|
||||||
|
maps.Copy(env, tc.env)
|
||||||
|
s, _, _ := startWithLookups(t, env)
|
||||||
|
|
||||||
|
// A quarter of 4 requests a minute: the second breaks the limit.
|
||||||
|
s.get(tc.from, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
line := s.get(tc.from, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||||
|
tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnknownLimitPercentGivesEveryClientWithoutACountryItsPercentage(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithLookups(t, map[string]string{
|
||||||
|
rateLimitPerMinute: fourAMinute, unknownLimitPercent: "50",
|
||||||
|
})
|
||||||
|
|
||||||
|
// One the lookup database does not hold, and one on a private address,
|
||||||
|
// which is never looked up: the third request of each breaks half of 4.
|
||||||
|
for _, from := range []string{unplaced, "10.0.0.8"} {
|
||||||
|
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(from, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
}
|
||||||
|
|
||||||
|
// One in a country has the whole limit.
|
||||||
|
for range 3 {
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientWithoutAnAnswerInTimeHasTheUnknownLimitPercent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
|
||||||
|
// says, with a GeoJS that never answers.
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||||
|
time.Now, map[string]string{unknownLimitPercent: "0"})
|
||||||
|
|
||||||
|
// Once the second the request waits for its answer is up, the client
|
||||||
|
// counts as without a country, and its zero allowance refuses the
|
||||||
|
// request before it reaches the app.
|
||||||
|
serveFromDE(t, server, http.MethodGet, http.NoBody)
|
||||||
|
|
||||||
|
line := out.requestLine(t)
|
||||||
|
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||||
|
"0 from "+unknownLimitPercent)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestWaitsForItsLookupWhileABiasedThresholdIsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const timeout = 3 * time.Second
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting, value string
|
||||||
|
waits bool
|
||||||
|
}{
|
||||||
|
{asnLimitPercent, asnDEHalf, true},
|
||||||
|
{countryLimitPercent, countryDEHalf, true},
|
||||||
|
{asnBytesPercent, asnDEHalf, true},
|
||||||
|
{countryBytesPercent, countryDEHalf, true},
|
||||||
|
{unknownLimitPercent, "99", true},
|
||||||
|
// At 100, its default, it lowers no limit.
|
||||||
|
{unknownLimitPercent, "100", false},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting+"="+tc.value, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
|
||||||
|
// says, with a GeoJS that never answers.
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
// The request's body is over SWWAF_REQUEST_MAX_BYTES, so that it
|
||||||
|
// is refused after the checks, and never reaches the app.
|
||||||
|
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||||
|
time.Now, map[string]string{
|
||||||
|
lookupTimeout: timeout.String(), requestMaxBytes: "1",
|
||||||
|
tc.setting: tc.value,
|
||||||
|
})
|
||||||
|
began := time.Now()
|
||||||
|
|
||||||
|
serveFromDE(t, server, http.MethodPost, strings.NewReader("ab"))
|
||||||
|
|
||||||
|
want := time.Duration(0)
|
||||||
|
if tc.waits {
|
||||||
|
want = timeout
|
||||||
|
}
|
||||||
|
|
||||||
|
if waited := time.Since(began); waited != want {
|
||||||
|
t.Errorf("the request waited %s for its answer, want %s", waited, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantLine(t, out.requestLine(t), http.StatusRequestEntityTooLarge,
|
||||||
|
requestlog.ActionTooLarge)
|
||||||
|
|
||||||
|
// The bubble's clock stops once this function returns, so the
|
||||||
|
// request to GeoJS, which a request that did not wait leaves
|
||||||
|
// under way, has to be abandoned before then.
|
||||||
|
time.Sleep(timeout)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBanForALoweredLimitGivesThePercentageInItsNotesAndItsAlert(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
env map[string]string
|
||||||
|
// before is how many uploads come before the one that breaks a
|
||||||
|
// limit, which is answered with status and logged with action.
|
||||||
|
before int
|
||||||
|
status int
|
||||||
|
action string
|
||||||
|
// reason and want are the ban's reason, and its notes' limit
|
||||||
|
// percentage, as percentText gives it.
|
||||||
|
reason, want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
// A quarter of 12 requests a minute is 3: the fourth breaks it.
|
||||||
|
"a rate limit",
|
||||||
|
map[string]string{rateLimitPerMinute: "12", asnLimitPercent: asnDEQuarter},
|
||||||
|
3, http.StatusForbidden, requestlog.ActionRateLimited,
|
||||||
|
"requests per minute over the limit of 3", "25 from " + asnLimitPercent,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// The byte limits' percentage, not the rate limits'.
|
||||||
|
"a byte limit",
|
||||||
|
map[string]string{
|
||||||
|
bytesLimitPerMinute: twoUploads, asnLimitPercent: asnDEQuarter,
|
||||||
|
asnBytesPercent: asnDEHalf,
|
||||||
|
},
|
||||||
|
0, http.StatusOK, requestlog.ActionForward,
|
||||||
|
"bytes per minute over the limit of 99", "50 from " + asnBytesPercent,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, server, queue := startWithLookups(t, tc.env)
|
||||||
|
|
||||||
|
for range tc.before {
|
||||||
|
s.uploadFrom(fromDE)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.requestWithBody(http.MethodPost, fromDE, "/", uploadHeader, uploadBody,
|
||||||
|
tc.status, tc.action)
|
||||||
|
|
||||||
|
held := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
|
||||||
|
if len(held) != 1 {
|
||||||
|
t.Fatalf("bans %+v, want one", held)
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := held[0].Notes
|
||||||
|
if held[0].Reason != tc.reason {
|
||||||
|
t.Errorf("the ban's reason is %q, want %q", held[0].Reason, tc.reason)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantPercent(t, "the notes' limit_percent", notes.LimitPercent,
|
||||||
|
notes.LimitPercentSetting, tc.want)
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 {
|
||||||
|
t.Fatalf("%d alerts wait, want the ban's alone: %+v", len(waiting), waiting)
|
||||||
|
}
|
||||||
|
|
||||||
|
alerted, _ := waiting[0].Detail["notes"].(bans.Notes)
|
||||||
|
wantPercent(t, "the alert's notes' limit_percent", alerted.LimitPercent,
|
||||||
|
alerted.LimitPercentSetting, tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// startWithLookups is startAppWithAlerts in front of readAndAnswer, with
|
||||||
|
// the settings in env on top of clients looked up in a lookup database,
|
||||||
|
// which places fromDE and fromKP in the AS numbers and countries the
|
||||||
|
// stand-in for GeoJS gives them, noCountry in AS64500 and no country, and
|
||||||
|
// no other address. It returns the sender, the server and the queue of
|
||||||
|
// the alerts.
|
||||||
|
func startWithLookups(
|
||||||
|
t *testing.T, env map[string]string,
|
||||||
|
) (*sender, *proxy.Server, *alerts.Queue) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||||
|
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
|
||||||
|
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
|
||||||
|
noCountry + "/32": {ASN: "AS64500", ASName: "Nowhere Net"},
|
||||||
|
})
|
||||||
|
|
||||||
|
settings := map[string]string{lookupSource: fileSource, lookupDBPath: path}
|
||||||
|
maps.Copy(settings, env)
|
||||||
|
|
||||||
|
s, _, server, queue := startAppWithAlerts(t, readAndAnswer, settings)
|
||||||
|
|
||||||
|
return s, server, queue
|
||||||
|
}
|
||||||
|
|
||||||
|
// uploadFrom is upload from the client at from.
|
||||||
|
func (s *sender) uploadFrom(from string) logLine {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
line, _ := s.requestWithBody(http.MethodPost, from, "/", uploadHeader, uploadBody,
|
||||||
|
http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
return line
|
||||||
|
}
|
||||||
|
|
||||||
|
// serveFromDE hands a request from fromDE with method and body straight to
|
||||||
|
// server's handler, without the network, and returns once it is answered.
|
||||||
|
func serveFromDE(t *testing.T, server *proxy.Server, method string, body io.Reader) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
req := httptest.NewRequestWithContext(t.Context(), method, "/", body)
|
||||||
|
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
|
||||||
|
|
||||||
|
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantPercent checks a limit percentage that a log line or a ban's notes
|
||||||
|
// give, what, and the setting that gave it, against want, as percentText
|
||||||
|
// gives them.
|
||||||
|
func wantPercent(t *testing.T, what string, percent *int64, setting, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if got := percentText(percent, setting); got != want {
|
||||||
|
t.Errorf("%s is %s, want %s", what, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// percentText gives a limit percentage and the setting that gave it as
|
||||||
|
// text, such as "50 from SWWAF_ASN_LIMIT_PERCENT", or none when both are
|
||||||
|
// left out.
|
||||||
|
func percentText(percent *int64, setting string) string {
|
||||||
|
switch {
|
||||||
|
case percent == nil && setting == "":
|
||||||
|
return none
|
||||||
|
case percent == nil:
|
||||||
|
return "none from " + setting
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("%d from %s", *percent, setting)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -103,6 +103,48 @@ func (b *responseBody) Close() error {
|
|||||||
return b.body.Close()
|
return b.body.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// upgradedConn is the connection to the app once the app has switched
|
||||||
|
// protocols, as for a WebSocket. ReverseProxy writes to it what the client
|
||||||
|
// sends and reads from it what the app sends, on goroutines of its own,
|
||||||
|
// until the connection closes; it counts the bytes each way, for the byte
|
||||||
|
// limits.
|
||||||
|
type upgradedConn struct {
|
||||||
|
io.ReadWriteCloser
|
||||||
|
|
||||||
|
// fromApp is how many bytes the app has sent, and toApp how many the
|
||||||
|
// client has.
|
||||||
|
fromApp atomic.Int64
|
||||||
|
toApp atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read reads what the app sends.
|
||||||
|
func (c *upgradedConn) Read(p []byte) (int, error) {
|
||||||
|
n, err := c.ReadWriteCloser.Read(p)
|
||||||
|
c.fromApp.Add(int64(n))
|
||||||
|
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write sends the app what the client sent.
|
||||||
|
func (c *upgradedConn) Write(p []byte) (int, error) {
|
||||||
|
n, err := c.ReadWriteCloser.Write(p)
|
||||||
|
c.toApp.Add(int64(n))
|
||||||
|
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// CloseWrite tells the app that the client sends no more, while what the
|
||||||
|
// app sends still passes. ReverseProxy calls it once the client has
|
||||||
|
// stopped sending, and closes the connection there if it is not supported.
|
||||||
|
func (c *upgradedConn) CloseWrite() error {
|
||||||
|
conn, ok := c.ReadWriteCloser.(interface{ CloseWrite() error })
|
||||||
|
if !ok {
|
||||||
|
return http.ErrNotSupported
|
||||||
|
}
|
||||||
|
|
||||||
|
return conn.CloseWrite()
|
||||||
|
}
|
||||||
|
|
||||||
// limitBody returns body, cut off with an *http.MaxBytesError after
|
// limitBody returns body, cut off with an *http.MaxBytesError after
|
||||||
// maxBytes, or unchanged if maxBytes is zero, which is off.
|
// maxBytes, or unchanged if maxBytes is zero, which is off.
|
||||||
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
|
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
|
||||||
|
|||||||
@@ -0,0 +1,583 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The byte limit settings.
|
||||||
|
const (
|
||||||
|
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
|
||||||
|
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
|
||||||
|
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
|
||||||
|
bytesCount = "SWWAF_BYTES_COUNT"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The values of SWWAF_BYTES_COUNT.
|
||||||
|
const (
|
||||||
|
countResponse = "response"
|
||||||
|
countRequest = "request"
|
||||||
|
countBoth = "both"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// bodyBytes is the size of the body of each request these tests send
|
||||||
|
// with one, and answerBytes that of each answer of the app.
|
||||||
|
bodyBytes = 30
|
||||||
|
answerBytes = 70
|
||||||
|
// byteLimit is the byte limit these tests set, as a setting: a request
|
||||||
|
// with a body and its answer, 100 bytes, go over it.
|
||||||
|
byteLimit = "99"
|
||||||
|
// minuteBytes is limit_hit for SWWAF_BYTES_LIMIT_PER_MINUTE.
|
||||||
|
minuteBytes = "minute_bytes"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEachByteLimitBansOnceTheResponseHasEnded(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const scraper = "192.0.2.200"
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting, window string
|
||||||
|
// apart is the time between the two requests, which the window
|
||||||
|
// still covers.
|
||||||
|
apart time.Duration
|
||||||
|
}{
|
||||||
|
{bytesLimitPerMinute, minute, 0},
|
||||||
|
{bytesLimitPerHour, "hour", 2 * time.Minute},
|
||||||
|
{bytesLimitPerDay, "day", 2 * time.Hour},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk := startWithAnswers(t, map[string]string{
|
||||||
|
tc.setting: byteLimit, metricsToken: token,
|
||||||
|
})
|
||||||
|
|
||||||
|
// 70 bytes are within the limit of 99.
|
||||||
|
line, _ := s.download()
|
||||||
|
if line.LimitHit != "" || line.Offence != "" {
|
||||||
|
t.Errorf("log line has limit_hit %q and offence %q, want neither",
|
||||||
|
line.LimitHit, line.Offence)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 140 bytes are over it. The response is passed on whole, and
|
||||||
|
// then bans the client for an hour.
|
||||||
|
clk.advance(tc.apart)
|
||||||
|
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
|
||||||
|
|
||||||
|
line, got := s.download()
|
||||||
|
if got.err != nil || len(got.body) != answerBytes ||
|
||||||
|
line.ResponseBytes != answerBytes {
|
||||||
|
t.Errorf("got %d bytes (%v), and the log line has response_bytes %d, "+
|
||||||
|
"want %d", len(got.body), got.err, line.ResponseBytes, answerBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
if line.LimitHit != tc.window+"_bytes" || line.Offence != requestlog.OffenceLimit ||
|
||||||
|
line.BanExpires != expires {
|
||||||
|
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
|
||||||
|
"want %s_bytes, limit and %s", line.LimitHit, line.Offence,
|
||||||
|
line.BanExpires, tc.window, expires)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
wantMetric(t, s.scrape(scraper), `smallwebwaf_rate_limit_hits_total{`+
|
||||||
|
`instance="`+alertInstance+`",kind="bytes",window="`+tc.window+`"}`, 1)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResponseOverAByteLimitByItselfIsPassedOnWhole(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _ := startWithAnswers(t, map[string]string{bytesLimitPerMinute: "50"})
|
||||||
|
|
||||||
|
// The answer's 70 bytes are over the limit of 50 on their own.
|
||||||
|
line, got := s.download()
|
||||||
|
if got.err != nil || len(got.body) != answerBytes || line.LimitHit != minuteBytes {
|
||||||
|
t.Errorf("got %d bytes (%v), and the log line has limit_hit %q, want %d and %s",
|
||||||
|
len(got.body), got.err, line.LimitHit, answerBytes, minuteBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBytesOfAnAnswerThatBreaksOffAreCounted(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, _, _ := startAppWithAlerts(t, breakOff, map[string]string{
|
||||||
|
bytesLimitPerMinute: "50",
|
||||||
|
})
|
||||||
|
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
|
||||||
|
|
||||||
|
// The 70 bytes passed on before the app broke off are over the limit of
|
||||||
|
// 50, and ban the client for an hour.
|
||||||
|
line, got := s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
|
||||||
|
requestlog.ActionUpstreamError)
|
||||||
|
if len(got.body) != answerBytes || line.LimitHit != minuteBytes ||
|
||||||
|
line.BanExpires != expires {
|
||||||
|
t.Errorf("got %d bytes, and the log line has limit_hit %q and ban_expires %q, "+
|
||||||
|
"want %d, %s and %s", len(got.body), line.LimitHit, line.BanExpires,
|
||||||
|
answerBytes, minuteBytes, expires)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebSocketBytesAreCountedOnceItCloses(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting string
|
||||||
|
counted float64
|
||||||
|
}{
|
||||||
|
{countResponse, answerBytes},
|
||||||
|
{countRequest, bodyBytes},
|
||||||
|
{countBoth, bodyBytes + answerBytes},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _, _ := startAppWithAlerts(t, answerAfterUpgrade, map[string]string{
|
||||||
|
bytesLimitPerMinute: "29", bytesCount: tc.setting,
|
||||||
|
})
|
||||||
|
|
||||||
|
// The client sends 30 bytes and the app 70, each over the limit
|
||||||
|
// of 29, which bans the client once the WebSocket has closed.
|
||||||
|
line := s.webSocket()
|
||||||
|
if line.LimitHit != minuteBytes || line.Counts.MinuteBytes != tc.counted {
|
||||||
|
t.Errorf("log line has limit_hit %q and minute_bytes %v, want %s and %v",
|
||||||
|
line.LimitHit, line.Counts.MinuteBytes, minuteBytes, tc.counted)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWebSocketPassesTheAnswerAfterTheClientStopsSending(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, echoOnceTheClientStops)
|
||||||
|
addr, out := startProxy(t, app.URL,
|
||||||
|
map[string]string{trustedProxies: trustLocalhost})
|
||||||
|
s := &sender{t: t, addr: addr, out: out}
|
||||||
|
|
||||||
|
conn, reader := s.openWebSocket()
|
||||||
|
send(t, conn, uploadBody)
|
||||||
|
|
||||||
|
// The client closes its sending side and waits for the answer, which the
|
||||||
|
// app sends only once it has seen the client stop. smallwebwaf passes the
|
||||||
|
// close on to the app through CloseWrite on upgradedConn; without that,
|
||||||
|
// it closes both connections, and the answer is lost.
|
||||||
|
tcp, ok := conn.(*net.TCPConn)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("connection is a %T, want a *net.TCPConn", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
err := tcp.CloseWrite()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("close the sending side: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, err := io.ReadAll(reader)
|
||||||
|
if err != nil || string(got) != uploadBody {
|
||||||
|
t.Errorf("got %q (%v), want %q", got, err, uploadBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.closeWebSocket(conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBytesCountSaysWhichBytesCount(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting string
|
||||||
|
// each is the bytes each request counts, and breaking the request
|
||||||
|
// that goes over the limit of 99.
|
||||||
|
each float64
|
||||||
|
breaking int
|
||||||
|
}{
|
||||||
|
{countResponse, answerBytes, 2},
|
||||||
|
{countRequest, bodyBytes, 4},
|
||||||
|
{countBoth, bodyBytes + answerBytes, 1},
|
||||||
|
} {
|
||||||
|
t.Run(tc.setting, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _ := startWithAnswers(t, map[string]string{
|
||||||
|
bytesLimitPerMinute: byteLimit, bytesCount: tc.setting,
|
||||||
|
})
|
||||||
|
|
||||||
|
for i := 1; i <= tc.breaking; i++ {
|
||||||
|
line := s.upload()
|
||||||
|
|
||||||
|
want := ""
|
||||||
|
if i == tc.breaking {
|
||||||
|
want = minuteBytes
|
||||||
|
}
|
||||||
|
|
||||||
|
counted := float64(i) * tc.each
|
||||||
|
if line.LimitHit != want || line.Counts.MinuteBytes != counted {
|
||||||
|
t.Errorf("request %d: log line has limit_hit %q and minute_bytes %v, "+
|
||||||
|
"want %q and %v", i, line.LimitHit, line.Counts.MinuteBytes,
|
||||||
|
want, counted)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestByteLimitsLeaveOutWhatTheRateLimitsLeaveOut(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
|
||||||
|
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||||
|
)
|
||||||
|
|
||||||
|
s, _ := startWithAnswers(t, map[string]string{
|
||||||
|
bytesLimitPerMinute: byteLimit,
|
||||||
|
allowNets: allowed,
|
||||||
|
rateLimitExemptNets: exempt,
|
||||||
|
rateLimitExemptPaths: "/assets/",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Each sends 200 bytes, none of which is counted.
|
||||||
|
for _, sent := range []struct{ from, path string }{
|
||||||
|
{allowed, "/"}, {exempt, "/"}, {client, "/assets/app.js"},
|
||||||
|
} {
|
||||||
|
for range 2 {
|
||||||
|
line, _ := s.requestWithBody(http.MethodPost, sent.from, sent.path,
|
||||||
|
uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||||
|
if _, counted := line.fields["counts"]; counted || line.LimitHit != "" {
|
||||||
|
t.Errorf("%s %s: log line has counts %v and limit_hit %q, want neither",
|
||||||
|
sent.from, sent.path, line.fields["counts"], line.LimitHit)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A path that is not exempt is counted, and breaks the limit.
|
||||||
|
line := s.upload()
|
||||||
|
if line.LimitHit != minuteBytes {
|
||||||
|
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestByteLimitsOffCountTheBytesAndBanNoOne(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const off = "off"
|
||||||
|
|
||||||
|
s, _ := startWithAnswers(t, map[string]string{
|
||||||
|
bytesLimitPerMinute: off, bytesLimitPerHour: off, bytesLimitPerDay: off,
|
||||||
|
})
|
||||||
|
|
||||||
|
for i := 1; i <= 3; i++ {
|
||||||
|
line := s.upload()
|
||||||
|
|
||||||
|
counted := float64(i * (bodyBytes + answerBytes))
|
||||||
|
if line.LimitHit != "" || line.Counts.MinuteBytes != counted ||
|
||||||
|
line.Counts.HourBytes != counted || line.Counts.DayBytes != counted {
|
||||||
|
t.Errorf("request %d: log line has limit_hit %q and counts %+v, "+
|
||||||
|
"want none and %v bytes in each window", i, line.LimitHit,
|
||||||
|
line.Counts, counted)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBanForABrokenByteLimitHasItsNotesAndItsAlert(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||||
|
bytesLimitPerMinute: byteLimit,
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
s.requestWithBody(http.MethodPost, client, "/upload?part=1", uploadHeader,
|
||||||
|
uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
netblock := netip.MustParsePrefix(client + "/32")
|
||||||
|
want := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: start,
|
||||||
|
Expires: start.Add(time.Hour),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
Reason: "bytes per minute over the limit of " + byteLimit,
|
||||||
|
Notes: bans.Notes{
|
||||||
|
Kind: "bytes",
|
||||||
|
Limit: 99,
|
||||||
|
Window: minute,
|
||||||
|
Count: bodyBytes + answerBytes,
|
||||||
|
// The request as it was answered, by the app.
|
||||||
|
Request: bans.Request{
|
||||||
|
Time: start,
|
||||||
|
Method: http.MethodPost,
|
||||||
|
Host: appHost,
|
||||||
|
Path: "/upload?part=1",
|
||||||
|
Status: http.StatusOK,
|
||||||
|
UserAgent: userAgent,
|
||||||
|
},
|
||||||
|
Requests: 1,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := server.Ledger.Bans(netblock)
|
||||||
|
if len(got) != 1 || got[0] != want {
|
||||||
|
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, want,
|
||||||
|
requestlog.FormatTime(want.Expires)))
|
||||||
|
|
||||||
|
if offences := historyOf(t, server, client).Offences.Limit; offences != 1 {
|
||||||
|
t.Errorf("history counts %d offences for a limit, want 1", offences)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeLogsAndAlertsAByteLimitAndBansNoOne(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
bytesLimitPerMinute: byteLimit,
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
// No ban sets the client's counters back to zero, so each request
|
||||||
|
// breaks the limit again. The answer is the app's either way, and the
|
||||||
|
// alert for the ban is not sent twice within the cooldown.
|
||||||
|
for range 2 {
|
||||||
|
line := s.upload()
|
||||||
|
wantWouldAction(t, line, "")
|
||||||
|
|
||||||
|
if line.LimitHit != minuteBytes || line.Offence != requestlog.OffenceLimit ||
|
||||||
|
line.BanExpires != "" {
|
||||||
|
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
|
||||||
|
"want %s, limit and none", line.LimitHit, line.Offence, line.BanExpires,
|
||||||
|
minuteBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
||||||
|
t.Errorf("the ledger holds %+v, want no ban", held)
|
||||||
|
}
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 {
|
||||||
|
t.Fatalf("%d alerts wait, want 1: %+v", len(waiting), waiting)
|
||||||
|
}
|
||||||
|
|
||||||
|
notes, _ := waiting[0].Detail["notes"].(bans.Notes)
|
||||||
|
alert := banAlert(alerts.EventBan, start, client, bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
|
||||||
|
Reason: "bytes per minute over the limit of " + byteLimit, Notes: notes,
|
||||||
|
}, requestlog.FormatTime(start.Add(time.Hour)))
|
||||||
|
alert.Detail["mode"] = observe
|
||||||
|
wantAlerts(t, queue, alert)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeLeavesOutTheBytesOfARequestEnforceModeRefuses(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _ := startWithAnswers(t, map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
bytesLimitPerMinute: "150",
|
||||||
|
})
|
||||||
|
|
||||||
|
s.upload()
|
||||||
|
|
||||||
|
// The second request breaks the rate limit, which in enforce mode would
|
||||||
|
// refuse it before the app sent anything, so its 100 bytes are not
|
||||||
|
// counted, and the byte limit is not broken. Its line gives the bytes
|
||||||
|
// counted before it.
|
||||||
|
line := s.upload()
|
||||||
|
wantWouldAction(t, line, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
if line.LimitHit != minute || line.Counts.MinuteBytes != bodyBytes+answerBytes {
|
||||||
|
t.Errorf("log line has limit_hit %q and minute_bytes %v, want minute and %d",
|
||||||
|
line.LimitHit, line.Counts.MinuteBytes, bodyBytes+answerBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// uploadHeader and uploadBody are the header and the body of a request
|
||||||
|
// with a body of bodyBytes.
|
||||||
|
//
|
||||||
|
//nolint:gochecknoglobals // a constant cannot call strings.Repeat
|
||||||
|
var (
|
||||||
|
uploadHeader = "Content-Length: " + strconv.Itoa(bodyBytes)
|
||||||
|
uploadBody = strings.Repeat("u", bodyBytes)
|
||||||
|
)
|
||||||
|
|
||||||
|
// readAndAnswer is the app of these tests: it reads each request's whole
|
||||||
|
// body and answers with answerBytes bytes.
|
||||||
|
func readAndAnswer(w http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
|
||||||
|
}
|
||||||
|
|
||||||
|
// breakOff is an app that announces an answer of twice answerBytes, and
|
||||||
|
// breaks off after answerBytes.
|
||||||
|
func breakOff(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.Header().Set("Content-Length", strconv.Itoa(2*answerBytes))
|
||||||
|
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
|
||||||
|
}
|
||||||
|
|
||||||
|
// answerAfterUpgrade is an app that switches protocols, as for a
|
||||||
|
// WebSocket, and then answers each line it receives with a line of
|
||||||
|
// answerBytes.
|
||||||
|
func answerAfterUpgrade(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
conn, buffered, err := http.NewResponseController(w).Hijack()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = conn.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
|
||||||
|
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
|
_ = buffered.Flush()
|
||||||
|
|
||||||
|
for {
|
||||||
|
_, err := buffered.ReadString('\n')
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _ = buffered.WriteString(strings.Repeat("a", answerBytes-1) + "\n")
|
||||||
|
_ = buffered.Flush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// echoOnceTheClientStops is an app that switches protocols, as for a
|
||||||
|
// WebSocket, reads what the client sends until the client stops sending,
|
||||||
|
// and then sends it all back.
|
||||||
|
func echoOnceTheClientStops(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
conn, buffered, err := http.NewResponseController(w).Hijack()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = conn.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
|
||||||
|
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
|
_ = buffered.Flush()
|
||||||
|
|
||||||
|
received, _ := io.ReadAll(buffered)
|
||||||
|
_, _ = buffered.Write(received)
|
||||||
|
_ = buffered.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// webSocket opens a WebSocket from client to answerAfterUpgrade, sends a
|
||||||
|
// line of bodyBytes on it, reads the answer, and closes it. It checks the
|
||||||
|
// answer, and the log line as request does, and returns the log line.
|
||||||
|
func (s *sender) webSocket() logLine {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
conn, reader := s.openWebSocket()
|
||||||
|
send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n")
|
||||||
|
|
||||||
|
got, err := reader.ReadString('\n')
|
||||||
|
if err != nil || len(got) != answerBytes {
|
||||||
|
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.closeWebSocket(conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// openWebSocket sends a request from client to switch protocols, as for a
|
||||||
|
// WebSocket, and checks that the app switches. It returns the connection,
|
||||||
|
// on which reading fails once waitLimit has passed, and a reader of what
|
||||||
|
// the app sends on it.
|
||||||
|
func (s *sender) openWebSocket() (net.Conn, *bufio.Reader) {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
conn := dial(s.t, s.addr)
|
||||||
|
send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+
|
||||||
|
": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
|
|
||||||
|
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
|
||||||
|
if err != nil {
|
||||||
|
s.t.Fatalf("set read deadline: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := bufio.NewReader(conn)
|
||||||
|
|
||||||
|
res, err := http.ReadResponse(reader, nil)
|
||||||
|
if err != nil {
|
||||||
|
s.t.Fatalf("read the answer to the upgrade: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = res.Body.Close()
|
||||||
|
|
||||||
|
if res.StatusCode != http.StatusSwitchingProtocols {
|
||||||
|
s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
||||||
|
}
|
||||||
|
|
||||||
|
return conn, reader
|
||||||
|
}
|
||||||
|
|
||||||
|
// closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its
|
||||||
|
// log line as request does, and returns it.
|
||||||
|
func (s *sender) closeWebSocket(conn net.Conn) logLine {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
|
||||||
|
s.sent++
|
||||||
|
wantLine(s.t, line, http.StatusSwitchingProtocols, requestlog.ActionForward)
|
||||||
|
|
||||||
|
return line
|
||||||
|
}
|
||||||
|
|
||||||
|
// startWithAnswers is startAppWithAlerts in front of readAndAnswer, for a
|
||||||
|
// test that looks at neither the server nor the alerts.
|
||||||
|
func startWithAnswers(t *testing.T, env map[string]string) (*sender, *clock) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
s, clk, _, _ := startAppWithAlerts(t, readAndAnswer, env)
|
||||||
|
|
||||||
|
return s, clk
|
||||||
|
}
|
||||||
|
|
||||||
|
// download sends a GET request for / from client, and checks that the
|
||||||
|
// app's answer is passed on, as request does. It returns the log line and
|
||||||
|
// the answer.
|
||||||
|
func (s *sender) download() (logLine, answer) {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
return s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
|
||||||
|
requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
// upload is download for a POST request with a body of bodyBytes, and
|
||||||
|
// returns the log line.
|
||||||
|
func (s *sender) upload() logLine {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
line, _ := s.requestWithBody(http.MethodPost, client, "/", uploadHeader,
|
||||||
|
uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
return line
|
||||||
|
}
|
||||||
@@ -1,31 +1,17 @@
|
|||||||
package proxy
|
package proxy
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
"slices"
|
||||||
)
|
)
|
||||||
|
|
||||||
// countryDenied reports whether the country lists refuse the request.
|
// countryDenied reports whether the country lists refuse the request, by
|
||||||
// The client's country is looked up only while a list is set, and never
|
// the client's country as it was looked up. A client without a country,
|
||||||
// for a client on a private, loopback or link-local address, which has
|
// or whose country cannot be found, is refused only by
|
||||||
// no country. A client without a country, or whose country cannot be
|
// SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES.
|
||||||
// found, is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. ctx is
|
func (rq *request) countryDenied() bool {
|
||||||
// the request's own context.
|
|
||||||
func (rq *request) countryDenied(ctx context.Context) bool {
|
|
||||||
denied := rq.h.config.DeniedCountries
|
denied := rq.h.config.DeniedCountries
|
||||||
allowed := rq.h.config.ExclusivelyAllowedCountries
|
allowed := rq.h.config.ExclusivelyAllowedCountries
|
||||||
|
country := rq.line.Country
|
||||||
if len(denied) == 0 && len(allowed) == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
var country string
|
|
||||||
if hasCountry(rq.client) {
|
|
||||||
country = rq.h.geojs.Country(ctx, clientGroup(rq.client))
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.line.Country = country
|
|
||||||
|
|
||||||
if slices.Contains(denied, country) {
|
if slices.Contains(denied, country) {
|
||||||
return true
|
return true
|
||||||
@@ -33,9 +19,3 @@ func (rq *request) countryDenied(ctx context.Context) bool {
|
|||||||
|
|
||||||
return len(allowed) > 0 && !slices.Contains(allowed, country)
|
return len(allowed) > 0 && !slices.Contains(allowed, country)
|
||||||
}
|
}
|
||||||
|
|
||||||
// hasCountry reports whether addr can be placed in a country: private,
|
|
||||||
// loopback and link-local addresses cannot.
|
|
||||||
func hasCountry(addr netip.Addr) bool {
|
|
||||||
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -56,16 +56,21 @@ func TestCountryLists(t *testing.T) {
|
|||||||
maps.Copy(env, tc.env)
|
maps.Copy(env, tc.env)
|
||||||
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
|
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
|
||||||
|
|
||||||
for i, sent := range []struct{ client, country string }{
|
// The AS number GeoJS gives unplaced, 64512, counts as unknown.
|
||||||
{fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""},
|
for i, sent := range []struct{ client, asn, asName, country string }{
|
||||||
|
{fromDE, asnDE, asNameDE, "DE"}, {fromKP, asnKP, asNameKP, "KP"},
|
||||||
|
{unplaced, "", "", ""},
|
||||||
} {
|
} {
|
||||||
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
req.Header.Set(forwardedFor, sent.client)
|
req.Header.Set(forwardedFor, sent.client)
|
||||||
got := do(t, req)
|
got := do(t, req)
|
||||||
|
|
||||||
line := out.requestLines(t, i+1)[i]
|
line := out.requestLines(t, i+1)[i]
|
||||||
if line.Country != sent.country {
|
if line.ASN != sent.asn || line.ASName != sent.asName ||
|
||||||
t.Errorf("log line has country %q, want %q", line.Country, sent.country)
|
line.Country != sent.country {
|
||||||
|
t.Errorf("log line has %q, %q and %q, want %q, %q and %q",
|
||||||
|
line.ASN, line.ASName, line.Country,
|
||||||
|
sent.asn, sent.asName, sent.country)
|
||||||
}
|
}
|
||||||
|
|
||||||
if slices.Contains(tc.refused, sent.client) {
|
if slices.Contains(tc.refused, sent.client) {
|
||||||
@@ -130,7 +135,7 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}}
|
answer := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "DE"}}
|
||||||
|
|
||||||
err := json.NewEncoder(w).Encode(answer)
|
err := json.NewEncoder(w).Encode(answer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -174,20 +179,15 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
|
|||||||
wantStatus(t, got, http.StatusOK)
|
wantStatus(t, got, http.StatusOK)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
|
func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
for _, tc := range []struct {
|
for _, tc := range []struct {
|
||||||
name string
|
name string
|
||||||
env map[string]string
|
env map[string]string
|
||||||
clients []string // "" sends no X-Forwarded-For: the client is 127.0.0.1
|
|
||||||
}{
|
}{
|
||||||
{"no country list is set", nil, []string{fromKP, fromDE}},
|
{"no setting needs the lookup", nil},
|
||||||
{
|
{"a country list is set", map[string]string{deniedCountries: "kp"}},
|
||||||
"private, loopback and link-local addresses",
|
|
||||||
map[string]string{deniedCountries: "kp"},
|
|
||||||
[]string{"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9"},
|
|
||||||
},
|
|
||||||
} {
|
} {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -198,7 +198,10 @@ func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
|
|||||||
maps.Copy(env, tc.env)
|
maps.Copy(env, tc.env)
|
||||||
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
|
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
|
||||||
|
|
||||||
for i, sent := range tc.clients {
|
// "" sends no X-Forwarded-For: the client is 127.0.0.1.
|
||||||
|
for i, sent := range []string{
|
||||||
|
"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9",
|
||||||
|
} {
|
||||||
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
if sent != "" {
|
if sent != "" {
|
||||||
req.Header.Set(forwardedFor, sent)
|
req.Header.Set(forwardedFor, sent)
|
||||||
@@ -209,15 +212,26 @@ func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
|
|||||||
line := out.requestLines(t, i+1)[i]
|
line := out.requestLines(t, i+1)[i]
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
country, present := line.fields["country"]
|
for _, field := range []string{"asn", "as_name", "country"} {
|
||||||
if !present || country != "" {
|
value, present := line.fields[field]
|
||||||
t.Errorf("log line for %q has country %v, want an empty one",
|
if !present || value != "" {
|
||||||
line.ClientIP, country)
|
t.Errorf("log line for %q has %s %v, want an empty one",
|
||||||
|
line.ClientIP, field, value)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(asked()) != 0 {
|
// GeoJS is asked about up to 200 waiting clients at once, so once it
|
||||||
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
// has been asked about fromDE, which comes last, it has been asked
|
||||||
|
// about every client before it that waited for an answer.
|
||||||
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
|
req.Header.Set(forwardedFor, fromDE)
|
||||||
|
wantStatus(t, do(t, req), http.StatusOK)
|
||||||
|
|
||||||
|
waitUntil(func() bool { return slices.Contains(asked(), fromDE) })
|
||||||
|
|
||||||
|
if got := asked(); !slices.Equal(got, []string{fromDE}) {
|
||||||
|
t.Errorf("GeoJS was asked about %v, want %s alone", got, fromDE)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -264,18 +278,31 @@ func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP
|
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP,
|
||||||
// and no other address. It returns its URL, and what returns the
|
// each in an AS of its own, and no other address. It returns its URL, and
|
||||||
// addresses it has been asked about.
|
// what returns the addresses it has been asked about.
|
||||||
func startGeoJS(t *testing.T) (string, func() []string) {
|
func startGeoJS(t *testing.T) (string, func() []string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
places := map[string]string{fromDE: "DE", fromKP: "KP"}
|
geojsURL, asked, release := startHeldGeoJS(t)
|
||||||
|
release()
|
||||||
|
|
||||||
var asked struct {
|
return geojsURL, asked
|
||||||
mu sync.Mutex
|
}
|
||||||
addrs []string
|
|
||||||
}
|
// startHeldGeoJS is startGeoJS for a stand-in that answers nothing until
|
||||||
|
// release is called. Each request to it waits until then.
|
||||||
|
func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
var (
|
||||||
|
asked struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
addrs []string
|
||||||
|
}
|
||||||
|
released = make(chan struct{})
|
||||||
|
once sync.Once
|
||||||
|
)
|
||||||
|
|
||||||
geojs := httptest.NewServer(http.HandlerFunc(
|
geojs := httptest.NewServer(http.HandlerFunc(
|
||||||
func(w http.ResponseWriter, r *http.Request) {
|
func(w http.ResponseWriter, r *http.Request) {
|
||||||
@@ -285,11 +312,11 @@ func startGeoJS(t *testing.T) (string, func() []string) {
|
|||||||
asked.addrs = append(asked.addrs, addrs...)
|
asked.addrs = append(asked.addrs, addrs...)
|
||||||
asked.mu.Unlock()
|
asked.mu.Unlock()
|
||||||
|
|
||||||
answers := make([]map[string]string, 0, len(addrs))
|
<-released
|
||||||
|
|
||||||
|
answers := make([]geojsAnswer, 0, len(addrs))
|
||||||
for _, addr := range addrs {
|
for _, addr := range addrs {
|
||||||
answers = append(answers, map[string]string{
|
answers = append(answers, answerAbout(addr))
|
||||||
"ip": addr, "country": places[addr],
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err := json.NewEncoder(w).Encode(answers)
|
err := json.NewEncoder(w).Encode(answers)
|
||||||
@@ -299,10 +326,48 @@ func startGeoJS(t *testing.T) (string, func() []string) {
|
|||||||
}))
|
}))
|
||||||
t.Cleanup(geojs.Close)
|
t.Cleanup(geojs.Close)
|
||||||
|
|
||||||
|
release := func() { once.Do(func() { close(released) }) }
|
||||||
|
// Run before geojs.Close, which waits for every request to be answered.
|
||||||
|
t.Cleanup(release)
|
||||||
|
|
||||||
return geojs.URL, func() []string {
|
return geojs.URL, func() []string {
|
||||||
asked.mu.Lock()
|
asked.mu.Lock()
|
||||||
defer asked.mu.Unlock()
|
defer asked.mu.Unlock()
|
||||||
|
|
||||||
return slices.Clone(asked.addrs)
|
return slices.Clone(asked.addrs)
|
||||||
|
}, release
|
||||||
|
}
|
||||||
|
|
||||||
|
// The AS numbers and names the stand-in for GeoJS gives fromDE and
|
||||||
|
// fromKP, as they are logged.
|
||||||
|
const (
|
||||||
|
asnDE = "AS64496"
|
||||||
|
asNameDE = "Example Net"
|
||||||
|
asnKP = "AS64511"
|
||||||
|
asNameKP = "Other Net"
|
||||||
|
)
|
||||||
|
|
||||||
|
// geojsAnswer is an answer of GeoJS about one address, with the fields
|
||||||
|
// smallwebwaf reads.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // GeoJS's own names
|
||||||
|
type geojsAnswer struct {
|
||||||
|
IP string `json:"ip"`
|
||||||
|
ASN int `json:"asn"`
|
||||||
|
ASName string `json:"organization_name"`
|
||||||
|
CountryCode string `json:"country_code,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// answerAbout is what the stand-in for GeoJS answers about addr: for an
|
||||||
|
// address it cannot place, the AS number 64512 and the AS name Unknown
|
||||||
|
// with no country, as GeoJS does.
|
||||||
|
func answerAbout(addr string) geojsAnswer {
|
||||||
|
switch addr {
|
||||||
|
case fromDE:
|
||||||
|
return geojsAnswer{IP: addr, ASN: 64496, ASName: asNameDE, CountryCode: "DE"}
|
||||||
|
case fromKP:
|
||||||
|
return geojsAnswer{IP: addr, ASN: 64511, ASName: asNameKP, CountryCode: "KP"}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return geojsAnswer{IP: addr, ASN: 64512, ASName: "Unknown"}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,7 +24,9 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
|
|||||||
start := clk.Now()
|
start := clk.Now()
|
||||||
|
|
||||||
// Two let through, one over the limit, which bans the client, and one
|
// Two let through, one over the limit, which bans the client, and one
|
||||||
// refused under that ban, for which the country is not looked up.
|
// refused under that ban, for which the client is not looked up. GeoJS
|
||||||
|
// answers about the client at its first request, and its later ones
|
||||||
|
// use that answer.
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
clk.advance(time.Second)
|
clk.advance(time.Second)
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
@@ -35,8 +37,10 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
|
|||||||
want := ratelimit.History{
|
want := ratelimit.History{
|
||||||
FirstSeen: start,
|
FirstSeen: start,
|
||||||
LastSeen: start.Add(2 * time.Second),
|
LastSeen: start.Add(2 * time.Second),
|
||||||
|
ASN: asnDE,
|
||||||
|
ASName: asNameDE,
|
||||||
Country: "DE",
|
Country: "DE",
|
||||||
LookedUp: start.Add(time.Second),
|
LookedUp: start,
|
||||||
Requests: 4,
|
Requests: 4,
|
||||||
Forwarded: 2,
|
Forwarded: 2,
|
||||||
Refused: 2,
|
Refused: 2,
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The headers in which the app is passed the client's AS number and
|
||||||
|
// country while SWWAF_ADD_LOOKUP_HEADERS is set. Go writes every header
|
||||||
|
// name in this form, as it sends it and as it receives it, so X-Client-ASN
|
||||||
|
// arrives as X-Client-Asn, and Del removes a client's own whatever their
|
||||||
|
// case; header names are not case-sensitive.
|
||||||
|
const (
|
||||||
|
asnHeader = "X-Client-Asn"
|
||||||
|
countryHeader = "X-Client-Country"
|
||||||
|
)
|
||||||
|
|
||||||
|
// lookUp looks up the client's AS number and country, in the lookup
|
||||||
|
// database or through GeoJS, and notes them for the log line, unless
|
||||||
|
// SWWAF_LOOKUP_SOURCE is off or the client is on a private, loopback or
|
||||||
|
// link-local address, which no lookup can place. The lookup database
|
||||||
|
// answers at once. With GeoJS, while a setting needs the answer, such as a
|
||||||
|
// country list or a biased threshold, a new client's request waits for it.
|
||||||
|
// ctx is the request's own context.
|
||||||
|
func (rq *request) lookUp(ctx context.Context) {
|
||||||
|
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if rq.h.config.LookupSource == "file" {
|
||||||
|
rq.lookupAnswer = rq.h.lookupFile.LookUp(clientGroup(rq.client))
|
||||||
|
} else {
|
||||||
|
rq.lookupAnswer = rq.h.geojs.LookUp(ctx, clientGroup(rq.client))
|
||||||
|
}
|
||||||
|
|
||||||
|
rq.lookedUp = true
|
||||||
|
rq.line.ASN = rq.lookupAnswer.ASN
|
||||||
|
rq.line.ASName = rq.lookupAnswer.ASName
|
||||||
|
rq.line.Country = rq.lookupAnswer.Country
|
||||||
|
}
|
||||||
|
|
||||||
|
// addLookup adds answer, an answer about a client from the lookup
|
||||||
|
// database or GeoJS, to the client's history, and to the notes of the bans
|
||||||
|
// on its netblock that have no AS number, AS name or country yet.
|
||||||
|
func (h *handler) addLookup(answer lookup.Answer) {
|
||||||
|
h.limiter.AddLookup(answer.Client, answer.Answered,
|
||||||
|
answer.ASN, answer.ASName, answer.Country)
|
||||||
|
h.ledger.AddLookup(h.netblock(answer.Client.Addr()),
|
||||||
|
answer.ASN, answer.ASName, answer.Country)
|
||||||
|
}
|
||||||
|
|
||||||
|
// setLookupHeaders sets the headers in which the app is passed the
|
||||||
|
// client's AS number and country, leaving out one that is unknown.
|
||||||
|
func setLookupHeaders(header http.Header, asn, country string) {
|
||||||
|
if asn != "" {
|
||||||
|
header.Set(asnHeader, asn)
|
||||||
|
}
|
||||||
|
|
||||||
|
if country != "" {
|
||||||
|
header.Set(countryHeader, country)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// canBePlaced reports whether a lookup can place addr: private, loopback
|
||||||
|
// and link-local addresses have no AS number or country.
|
||||||
|
func canBePlaced(addr netip.Addr) bool {
|
||||||
|
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
|
||||||
|
}
|
||||||
@@ -0,0 +1,377 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/netip"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// asnAndCountry is what a lookup gives a client: its AS number, AS name
|
||||||
|
// and country.
|
||||||
|
type asnAndCountry struct{ asn, asName, country string }
|
||||||
|
|
||||||
|
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
|
||||||
|
// lookup database.
|
||||||
|
const fileSource = "file"
|
||||||
|
|
||||||
|
func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// The stand-in for GeoJS answers only once released. A request that
|
||||||
|
// waited for it would wait an hour, and get no answer within
|
||||||
|
// waitLimit.
|
||||||
|
geojsURL, asked, release := startHeldGeoJS(t)
|
||||||
|
s, _, server := startWithClock(t, geojsURL, map[string]string{
|
||||||
|
lookupTimeout: "1h",
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
|
||||||
|
// fromDE's second request breaks the limit and bans it, and fromKP
|
||||||
|
// comes too. None waits for GeoJS.
|
||||||
|
for _, line := range []logLine{
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
|
||||||
|
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
|
||||||
|
s.get(fromKP, http.StatusOK, requestlog.ActionForward),
|
||||||
|
} {
|
||||||
|
got := asnAndCountry{line.ASN, line.ASName, line.Country}
|
||||||
|
if got != (asnAndCountry{}) {
|
||||||
|
t.Errorf("log line has %+v before GeoJS answered, want nothing", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Once GeoJS answers, each answer reaches the client's history, and
|
||||||
|
// fromDE's reaches the notes of its ban.
|
||||||
|
release()
|
||||||
|
|
||||||
|
netblock := netip.MustParsePrefix(fromDE + "/32")
|
||||||
|
|
||||||
|
waitUntil(func() bool {
|
||||||
|
return historyOf(t, server, fromDE).ASN != "" &&
|
||||||
|
historyOf(t, server, fromKP).ASN != "" &&
|
||||||
|
server.Ledger.Bans(netblock)[0].Notes.ASN != ""
|
||||||
|
})
|
||||||
|
|
||||||
|
de := asnAndCountry{asnDE, asNameDE, "DE"}
|
||||||
|
|
||||||
|
for addr, want := range map[string]asnAndCountry{
|
||||||
|
fromDE: de, fromKP: {asnKP, asNameKP, "KP"},
|
||||||
|
} {
|
||||||
|
h := historyOf(t, server, addr)
|
||||||
|
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want {
|
||||||
|
t.Errorf("%s's history has %+v, want %+v", addr, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := server.Ledger.Bans(netblock)[0].Notes
|
||||||
|
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
|
||||||
|
t.Errorf("the ban's notes have %+v, want %+v", got, de)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GeoJS was asked about each client once, fromKP after fromDE, whose
|
||||||
|
// request was under way when fromKP came.
|
||||||
|
if got := asked(); !slices.Equal(got, []string{fromDE, fromKP}) {
|
||||||
|
t.Errorf("GeoJS was asked about %v, want %s and %s", got, fromDE, fromKP)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestASNumberAndNameInTheLogLineTheHistoryTheBanNotesAndTheAlert(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
|
||||||
|
addr, out, server, queue := startProxyWithAlerts(t, app.URL, geojsURL, clk.Now,
|
||||||
|
map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
alertWebhookURL: "https://alerts.example/smallwebwaf",
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
s := &sender{t: t, addr: addr, out: out}
|
||||||
|
|
||||||
|
// The answer is kept before the requests, so GeoJS is not asked, and
|
||||||
|
// gives no answer of its own.
|
||||||
|
netblock := netip.MustParsePrefix(fromDE + "/32")
|
||||||
|
server.GeoJS.Load([]lookup.Answer{{
|
||||||
|
Client: netblock, ASN: asnDE, ASName: asNameDE, Country: "DE",
|
||||||
|
Answered: clk.Now(), Used: clk.Now(),
|
||||||
|
}})
|
||||||
|
|
||||||
|
want := asnAndCountry{asnDE, asNameDE, "DE"}
|
||||||
|
|
||||||
|
for _, line := range []logLine{
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
|
||||||
|
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
|
||||||
|
} {
|
||||||
|
if got := (asnAndCountry{line.ASN, line.ASName, line.Country}); got != want {
|
||||||
|
t.Errorf("log line has %+v, want %+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
h := historyOf(t, server, fromDE)
|
||||||
|
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want ||
|
||||||
|
!h.LookedUp.Equal(clk.Now()) {
|
||||||
|
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
|
||||||
|
got, h.LookedUp, want, clk.Now())
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := server.Ledger.Bans(netblock)[0].Notes
|
||||||
|
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != want {
|
||||||
|
t.Errorf("the ban's notes have %+v, want %+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 {
|
||||||
|
t.Fatalf("alerts waiting %+v, want the ban's alone", waiting)
|
||||||
|
}
|
||||||
|
|
||||||
|
alert := waiting[0]
|
||||||
|
if got := (asnAndCountry{alert.ASN, alert.ASName, alert.Country}); got != want {
|
||||||
|
t.Errorf("the ban's alert has %+v, want %+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupSourceOffLooksNoClientUp(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojsURL, asked := startGeoJS(t)
|
||||||
|
s, clk, server := startWithClock(t, geojsURL, map[string]string{lookupSource: "off"})
|
||||||
|
|
||||||
|
// Even an answer kept from before is not used.
|
||||||
|
server.GeoJS.Load([]lookup.Answer{{
|
||||||
|
Client: netip.MustParsePrefix(fromDE + "/32"), ASN: asnDE, ASName: asNameDE,
|
||||||
|
Country: "DE", Answered: clk.Now(), Used: clk.Now(),
|
||||||
|
}})
|
||||||
|
|
||||||
|
for _, from := range []string{fromDE, fromKP} {
|
||||||
|
line := s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
got := asnAndCountry{line.ASN, line.ASName, line.Country}
|
||||||
|
if got != (asnAndCountry{}) {
|
||||||
|
t.Errorf("log line has %+v, want nothing", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if h := historyOf(t, server, fromDE); h.ASN != "" || !h.LookedUp.IsZero() {
|
||||||
|
t.Errorf("history has %q, looked up at %s, want no lookup", h.ASN, h.LookedUp)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(asked()) != 0 {
|
||||||
|
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojsURL, asked := startGeoJS(t)
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||||
|
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
|
||||||
|
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
|
||||||
|
})
|
||||||
|
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
||||||
|
lookupSource: fileSource,
|
||||||
|
lookupDBPath: path,
|
||||||
|
allowedCountries: "DE",
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
|
||||||
|
// fromDE's second request breaks the limit and bans it. The list
|
||||||
|
// refuses fromKP, and unplaced, which the file does not hold.
|
||||||
|
de := asnAndCountry{asnDE, asNameDE, "DE"}
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
line logLine
|
||||||
|
want asnAndCountry
|
||||||
|
}{
|
||||||
|
{s.get(fromDE, http.StatusOK, requestlog.ActionForward), de},
|
||||||
|
{s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), de},
|
||||||
|
{
|
||||||
|
s.get(fromKP, http.StatusForbidden, requestlog.ActionCountryDenied),
|
||||||
|
asnAndCountry{asnKP, asNameKP, "KP"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
s.get(unplaced, http.StatusForbidden, requestlog.ActionCountryDenied),
|
||||||
|
asnAndCountry{},
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
got := asnAndCountry{tc.line.ASN, tc.line.ASName, tc.line.Country}
|
||||||
|
if got != tc.want {
|
||||||
|
t.Errorf("log line has %+v, want %+v", got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
h := historyOf(t, server, fromDE)
|
||||||
|
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != de ||
|
||||||
|
!h.LookedUp.Equal(clk.Now()) {
|
||||||
|
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
|
||||||
|
got, h.LookedUp, de, clk.Now())
|
||||||
|
}
|
||||||
|
|
||||||
|
notes := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))[0].Notes
|
||||||
|
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
|
||||||
|
t.Errorf("the ban's notes have %+v, want %+v", got, de)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(asked()) != 0 {
|
||||||
|
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var (
|
||||||
|
mu sync.Mutex
|
||||||
|
got [][2][]string // each request's X-Client-ASN and X-Client-Country
|
||||||
|
)
|
||||||
|
|
||||||
|
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
got = append(got, [2][]string{
|
||||||
|
r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country"),
|
||||||
|
})
|
||||||
|
})
|
||||||
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
addLookupHeaders: "true",
|
||||||
|
})
|
||||||
|
s := &sender{t: t, addr: addr, out: out}
|
||||||
|
|
||||||
|
// Each client sends headers of its own. fromDE's first request waits
|
||||||
|
// for its answer, which the app is passed; unplaced has none to pass,
|
||||||
|
// and a client on a private address is not looked up.
|
||||||
|
for _, from := range []string{fromDE, unplaced, "10.0.0.8"} {
|
||||||
|
s.requestWithHeader(from, "/", clientsOwnLookupHeaders,
|
||||||
|
http.StatusOK, requestlog.ActionForward)
|
||||||
|
}
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
want := [][2][]string{{{asnDE}, {"DE"}}, {nil, nil}, {nil, nil}}
|
||||||
|
if !slices.EqualFunc(got, want, func(a, b [2][]string) bool {
|
||||||
|
return slices.Equal(a[0], b[0]) && slices.Equal(a[1], b[1])
|
||||||
|
}) {
|
||||||
|
t.Errorf("the app was passed %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientsOwnLookupHeadersAreRemovedWhileTheSettingIsOff(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var (
|
||||||
|
mu sync.Mutex
|
||||||
|
asn, country []string
|
||||||
|
)
|
||||||
|
|
||||||
|
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
asn, country = r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country")
|
||||||
|
})
|
||||||
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
})
|
||||||
|
s := &sender{t: t, addr: addr, out: out}
|
||||||
|
|
||||||
|
s.requestWithHeader(fromDE, "/", clientsOwnLookupHeaders,
|
||||||
|
http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
if asn != nil || country != nil {
|
||||||
|
t.Errorf("the app was passed X-Client-ASN %v and X-Client-Country %v, want neither",
|
||||||
|
asn, country)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestWaitsAsLongAsTheLookupTimeoutSays(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// The test runs in a synctest bubble, where the time package runs on a
|
||||||
|
// clock of the test's own: the wait lasts exactly as long as it should,
|
||||||
|
// however slowly the test process runs. Nothing in it may wait on the
|
||||||
|
// network, which would keep that clock from moving on: the request is
|
||||||
|
// handed to the proxy's handler, and GeoJS is one that never answers.
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
// Not the default second. The exclusive list needs the answer, and
|
||||||
|
// the app is never reached.
|
||||||
|
const timeout = 3 * time.Second
|
||||||
|
|
||||||
|
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||||
|
time.Now, map[string]string{
|
||||||
|
lookupTimeout: timeout.String(),
|
||||||
|
allowedCountries: "DE",
|
||||||
|
})
|
||||||
|
|
||||||
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
|
||||||
|
http.NoBody)
|
||||||
|
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
|
||||||
|
began := time.Now()
|
||||||
|
|
||||||
|
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
|
||||||
|
|
||||||
|
if waited := time.Since(began); waited != timeout {
|
||||||
|
t.Errorf("the request waited %s for its answer, want %s", waited, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Without an answer, the client is in no country the list allows.
|
||||||
|
wantLine(t, out.requestLine(t), http.StatusForbidden,
|
||||||
|
requestlog.ActionCountryDenied)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// unansweredGeoJSURL is where a GeoJS that never answers is asked: a
|
||||||
|
// request to it waits, without the network, until it is abandoned.
|
||||||
|
// TestMain registers it with Go's default transport, through which GeoJS
|
||||||
|
// is asked.
|
||||||
|
const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json"
|
||||||
|
|
||||||
|
func TestMain(m *testing.M) {
|
||||||
|
transport, _ := http.DefaultTransport.(*http.Transport)
|
||||||
|
transport.RegisterProtocol("unanswered", unansweredGeoJS{})
|
||||||
|
|
||||||
|
m.Run()
|
||||||
|
}
|
||||||
|
|
||||||
|
// unansweredGeoJS is the GeoJS at unansweredGeoJSURL.
|
||||||
|
type unansweredGeoJS struct{}
|
||||||
|
|
||||||
|
// RoundTrip waits until req is abandoned.
|
||||||
|
func (unansweredGeoJS) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
<-req.Context().Done()
|
||||||
|
|
||||||
|
return nil, req.Context().Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// clientsOwnLookupHeaders are the X-Client-ASN and X-Client-Country a
|
||||||
|
// client sends of its own, each twice, in two cases.
|
||||||
|
const clientsOwnLookupHeaders = "X-Client-ASN: AS1\r\nx-client-asn: AS2\r\n" +
|
||||||
|
"X-CLIENT-COUNTRY: KP\r\nx-client-country: CN"
|
||||||
|
|
||||||
|
// waitUntil waits until done reports true, for at most waitLimit.
|
||||||
|
func waitUntil(done func() bool) {
|
||||||
|
deadline := time.Now().Add(waitLimit)
|
||||||
|
for !done() && time.Now().Before(deadline) {
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
+106
-45
@@ -148,8 +148,8 @@ func TestMetricsCountTheTraffic(t *testing.T) {
|
|||||||
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
|
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
|
||||||
out.requestLines(t, 2)
|
out.requestLines(t, 2)
|
||||||
|
|
||||||
forward := `{action="forward",status_class="2xx"}`
|
forward := `{action="forward",instance="app",status_class="2xx"}`
|
||||||
notFound := `{action="admin",status_class="4xx"}`
|
notFound := `{action="admin",instance="app",status_class="4xx"}`
|
||||||
|
|
||||||
// The request for the metrics is itself under way.
|
// The request for the metrics is itself under way.
|
||||||
metrics := scrape(t, addr)
|
metrics := scrape(t, addr)
|
||||||
@@ -159,11 +159,13 @@ func TestMetricsCountTheTraffic(t *testing.T) {
|
|||||||
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
|
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
|
||||||
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
|
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
|
||||||
float64(len("Not Found\n")))
|
float64(len("Not Found\n")))
|
||||||
wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
|
wantMetric(t, metrics,
|
||||||
wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
|
`smallwebwaf_request_duration_seconds_count{instance="app"}`, 2)
|
||||||
wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
|
wantMetric(t, metrics,
|
||||||
metric(t, metrics, "go_goroutines")
|
`smallwebwaf_upstream_duration_seconds_count{instance="app"}`, 1)
|
||||||
metric(t, metrics, "process_start_time_seconds")
|
wantMetric(t, metrics, `smallwebwaf_requests_in_flight{instance="app"}`, 1)
|
||||||
|
metric(t, metrics, `go_goroutines{instance="app"}`)
|
||||||
|
metric(t, metrics, `process_start_time_seconds{instance="app"}`)
|
||||||
|
|
||||||
// A request the app holds is under way until it ends.
|
// A request the app holds is under way until it ends.
|
||||||
httpClient := newClient(t)
|
httpClient := newClient(t)
|
||||||
@@ -180,7 +182,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
<-arrived
|
<-arrived
|
||||||
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
|
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 2)
|
||||||
releaseApp()
|
releaseApp()
|
||||||
|
|
||||||
err := <-ended
|
err := <-ended
|
||||||
@@ -189,7 +191,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
out.requestLines(t, 5)
|
out.requestLines(t, 5)
|
||||||
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
|
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMetricsCountLimitsAndBans(t *testing.T) {
|
func TestMetricsCountLimitsAndBans(t *testing.T) {
|
||||||
@@ -219,15 +221,16 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
|
|||||||
|
|
||||||
metrics := s.scrape(scraper)
|
metrics := s.scrape(scraper)
|
||||||
wantMetric(t, metrics,
|
wantMetric(t, metrics,
|
||||||
`smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
|
`smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1)
|
||||||
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
|
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
|
||||||
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
|
`kind="requests",window="minute"}`, 1)
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
|
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
|
||||||
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
|
||||||
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
|
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 0)
|
||||||
|
|
||||||
clk.advance(time.Hour)
|
clk.advance(time.Hour)
|
||||||
wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
|
wantMetric(t, s.scrape(scraper), `smallwebwaf_active_bans{instance="app"}`, 0)
|
||||||
|
|
||||||
// A limit broken again right after would ban for three hours, longer
|
// A limit broken again right after would ban for three hours, longer
|
||||||
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
|
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
|
||||||
@@ -235,13 +238,14 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
|
|||||||
s.get(client, 0, requestlog.ActionRateLimited)
|
s.get(client, 0, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
metrics = s.scrape(scraper)
|
metrics = s.scrape(scraper)
|
||||||
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
|
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
|
||||||
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
|
`kind="requests",window="minute"}`, 2)
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
|
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2)
|
||||||
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2)
|
||||||
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
|
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
|
||||||
// denied, client, and the scraper as of its earlier requests.
|
// denied, client, and the scraper as of its earlier requests.
|
||||||
wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
|
wantMetric(t, metrics, `smallwebwaf_tracked_clients{instance="app"}`, 3)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
|
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
|
||||||
@@ -255,7 +259,7 @@ func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
|
|||||||
rateLimitExemptNets: scraper,
|
rateLimitExemptNets: scraper,
|
||||||
})
|
})
|
||||||
|
|
||||||
const admins = `smallwebwaf_bans_made_total{cause="admin"}`
|
const admins = `smallwebwaf_bans_made_total{cause="admin",instance="app"}`
|
||||||
|
|
||||||
wantMetric(t, s.scrape(scraper), admins, 0)
|
wantMetric(t, s.scrape(scraper), admins, 0)
|
||||||
|
|
||||||
@@ -287,7 +291,8 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
|
|||||||
metricsTopN: "2",
|
metricsTopN: "2",
|
||||||
deniedCountries: "kp",
|
deniedCountries: "kp",
|
||||||
}
|
}
|
||||||
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, time.Now, env)
|
||||||
|
|
||||||
// The answers are kept before the requests, so that none waits for
|
// The answers are kept before the requests, so that none waits for
|
||||||
// GeoJS.
|
// GeoJS.
|
||||||
@@ -319,28 +324,82 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
|
|||||||
metrics := scrape(t, addr)
|
metrics := scrape(t, addr)
|
||||||
lines++
|
lines++
|
||||||
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
|
wantMetric(t, metrics,
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
|
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3)
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
|
wantMetric(t, metrics,
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
|
`smallwebwaf_country_requests_total{country="DE",instance="app"}`, 2)
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
|
wantMetric(t, metrics,
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
|
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 1)
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_list_refusals_total{country="KP",instance="app"}`, 3)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_request_bytes_total{country="KP",instance="app"}`, 0)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`, 6)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_response_bytes_total{country="KP",instance="app"}`,
|
||||||
float64(3*len("Forbidden\n")))
|
float64(3*len("Forbidden\n")))
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_response_bytes_total{country="other",instance="app"}`,
|
||||||
float64(len("hello")))
|
float64(len("hello")))
|
||||||
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
|
wantNoSeries(t, metrics,
|
||||||
|
`smallwebwaf_country_requests_total{country="FR",instance="app"}`)
|
||||||
|
|
||||||
// Once FR is busier than DE, it takes DE's place: its series counts
|
// Once FR is busier than DE, it takes DE's place: its series counts
|
||||||
// from then on, and DE's is gone.
|
// from then on, and DE's is gone.
|
||||||
send(fromFR, 3, http.StatusOK)
|
send(fromFR, 3, http.StatusOK)
|
||||||
|
|
||||||
metrics = scrape(t, addr)
|
metrics = scrape(t, addr)
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
|
wantMetric(t, metrics,
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
|
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3)
|
||||||
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
|
wantMetric(t, metrics,
|
||||||
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
|
`smallwebwaf_country_requests_total{country="FR",instance="app"}`, 2)
|
||||||
wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 2)
|
||||||
|
wantNoSeries(t, metrics,
|
||||||
|
`smallwebwaf_country_requests_total{country="DE",instance="app"}`)
|
||||||
|
wantNoSeries(t, metrics,
|
||||||
|
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsByASNumberKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
||||||
|
metricsToken: token,
|
||||||
|
metricsTopN: "1",
|
||||||
|
})
|
||||||
|
|
||||||
|
// The answers are kept before the requests, so that GeoJS gives none
|
||||||
|
// of its own. Each client is in an AS of its own.
|
||||||
|
answer := func(addr, asn string) lookup.Answer {
|
||||||
|
return lookup.Answer{
|
||||||
|
Client: netip.MustParsePrefix(addr + "/32"), ASN: asn,
|
||||||
|
Answered: clk.Now(), Used: clk.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
server.GeoJS.Load([]lookup.Answer{
|
||||||
|
answer(fromDE, "AS64501"), answer(fromKP, "AS64502"),
|
||||||
|
})
|
||||||
|
|
||||||
|
// With one AS number of its own, the other is counted as other. The
|
||||||
|
// metrics are asked for from a private address, which has no AS number.
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(fromKP, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
metrics := s.scrape("10.0.0.9")
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_asn_requests_total{asn="AS64501",instance="app"}`, 2)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_asn_requests_total{asn="other",instance="app"}`, 1)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_asn_request_bytes_total{asn="AS64501",instance="app"}`, 0)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_asn_response_bytes_total{asn="other",instance="app"}`, 0)
|
||||||
|
wantNoSeries(t, metrics,
|
||||||
|
`smallwebwaf_asn_requests_total{asn="AS64502",instance="app"}`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
|
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
|
||||||
@@ -370,16 +429,16 @@ func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
|
|||||||
deadline := time.Now().Add(waitLimit)
|
deadline := time.Now().Add(waitLimit)
|
||||||
metrics := scrape(t, addr)
|
metrics := scrape(t, addr)
|
||||||
|
|
||||||
for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
|
for metric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`) == 0 &&
|
||||||
time.Now().Before(deadline) {
|
time.Now().Before(deadline) {
|
||||||
time.Sleep(pollInterval)
|
time.Sleep(pollInterval)
|
||||||
|
|
||||||
metrics = scrape(t, addr)
|
metrics = scrape(t, addr)
|
||||||
}
|
}
|
||||||
|
|
||||||
wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
|
wantMetric(t, metrics, `smallwebwaf_geojs_requests_total{instance="app"}`, 1)
|
||||||
wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
|
wantMetric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`, 1)
|
||||||
wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
|
wantMetric(t, metrics, `smallwebwaf_geojs_unanswered_total{instance="app"}`, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
// keptAnswer returns GeoJS's answer that the client at addr is in
|
// keptAnswer returns GeoJS's answer that the client at addr is in
|
||||||
@@ -422,8 +481,9 @@ func (s *sender) scrape(from string) string {
|
|||||||
|
|
||||||
// metric returns the value of series in metrics, which are in the
|
// metric returns the value of series in metrics, which are in the
|
||||||
// Prometheus text format. series is a name and its labels in the order of
|
// Prometheus text format. series is a name and its labels in the order of
|
||||||
// their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
|
// their names, such as
|
||||||
// the test if there is no such series.
|
// smallwebwaf_offences_total{instance="app",kind="limit"}. It fails the
|
||||||
|
// test if there is no such series.
|
||||||
func metric(t *testing.T, metrics, series string) float64 {
|
func metric(t *testing.T, metrics, series string) float64 {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
@@ -471,7 +531,8 @@ func wantNoSeries(t *testing.T, metrics, series string) {
|
|||||||
func wantLimitHits(t *testing.T, addr, limit string, hits int) {
|
func wantLimitHits(t *testing.T, addr, limit string, hits int) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}`
|
series := `smallwebwaf_size_and_time_limit_hits_total{instance="app",limit="` +
|
||||||
|
limit + `"}`
|
||||||
metrics := scrape(t, addr)
|
metrics := scrape(t, addr)
|
||||||
|
|
||||||
if hits == 0 {
|
if hits == 0 {
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -123,10 +122,10 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
|
|||||||
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
hostname, _ := os.Hostname()
|
bytes := float64(sent + received)
|
||||||
|
|
||||||
want := withTimings(line, requestlog.Line{
|
want := withTimings(line, requestlog.Line{
|
||||||
Type: requestType, Time: line.Time, Instance: hostname,
|
Type: requestType, Time: line.Time, Instance: "app",
|
||||||
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
|
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
|
||||||
Path: rawPath, Query: rawQuery, Protocol: protocol,
|
Path: rawPath, Query: rawQuery, Protocol: protocol,
|
||||||
Status: http.StatusTeapot, RequestBytes: int64(sent),
|
Status: http.StatusTeapot, RequestBytes: int64(sent),
|
||||||
@@ -134,7 +133,10 @@ func wantRequestFields(t *testing.T, line logLine, host string, sent, received i
|
|||||||
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
|
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
|
||||||
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
|
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
|
||||||
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
|
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
|
||||||
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
Counts: ratelimit.Counts{
|
||||||
|
Minute: 1, Hour: 1, Day: 1,
|
||||||
|
MinuteBytes: bytes, HourBytes: bytes, DayBytes: bytes,
|
||||||
|
},
|
||||||
})
|
})
|
||||||
if !reflect.DeepEqual(line.Line, want) {
|
if !reflect.DeepEqual(line.Line, want) {
|
||||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||||
@@ -318,7 +320,7 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
|
|||||||
server := proxy.New(proxy.Params{
|
server := proxy.New(proxy.Params{
|
||||||
Config: cfg,
|
Config: cfg,
|
||||||
RequestLog: io.Discard,
|
RequestLog: io.Discard,
|
||||||
ProcessLog: requestlog.NewProcessLogger(io.Discard),
|
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName),
|
||||||
})
|
})
|
||||||
|
|
||||||
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
|
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
|
||||||
|
|||||||
+52
-22
@@ -11,6 +11,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
@@ -53,9 +54,12 @@ type Params struct {
|
|||||||
RequestLog io.Writer
|
RequestLog io.Writer
|
||||||
// ProcessLog receives the process's own messages.
|
// ProcessLog receives the process's own messages.
|
||||||
ProcessLog *slog.Logger
|
ProcessLog *slog.Logger
|
||||||
// GeoJSURL is where clients' countries are looked up, normally
|
// GeoJSURL is where clients' AS numbers and countries are looked up
|
||||||
// lookup.URL. GeoJS is asked only while a country list is set.
|
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL.
|
||||||
GeoJSURL string
|
GeoJSURL string
|
||||||
|
// LookupFile is the lookup database they are looked up in while
|
||||||
|
// SWWAF_LOOKUP_SOURCE is file, and nil otherwise.
|
||||||
|
LookupFile *lookup.File
|
||||||
// Now tells the time by which requests are counted for the rate
|
// Now tells the time by which requests are counted for the rate
|
||||||
// limits, bans are made and run out, and GeoJS's answers are kept,
|
// limits, bans are made and run out, and GeoJS's answers are kept,
|
||||||
// normally time.Now in UTC, the time the state files give.
|
// normally time.Now in UTC, the time the state files give.
|
||||||
@@ -63,17 +67,22 @@ type Params struct {
|
|||||||
// Rules are the rule files' rules, which each request is checked
|
// Rules are the rule files' rules, which each request is checked
|
||||||
// against.
|
// against.
|
||||||
Rules *rules.Files
|
Rules *rules.Files
|
||||||
|
// Alerts receive the alert for each ban the proxy makes or makes
|
||||||
|
// permanent, and for GeoJS failing.
|
||||||
|
Alerts *alerts.Queue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Server is the server smallwebwaf runs, with the parts of the proxy
|
// Server is the server smallwebwaf runs, with the parts of the proxy
|
||||||
// whose state the state files keep, and the metrics.
|
// whose state the state files keep, the lookup database, nil unless
|
||||||
|
// SWWAF_LOOKUP_SOURCE is file, and the metrics.
|
||||||
type Server struct {
|
type Server struct {
|
||||||
*http.Server
|
*http.Server
|
||||||
|
|
||||||
Ledger *bans.Ledger
|
Ledger *bans.Ledger
|
||||||
Limiter *ratelimit.Limiter
|
Limiter *ratelimit.Limiter
|
||||||
GeoJS *lookup.GeoJS
|
GeoJS *lookup.GeoJS
|
||||||
Metrics *metrics.Metrics
|
LookupFile *lookup.File
|
||||||
|
Metrics *metrics.Metrics
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns the server smallwebwaf runs: each request it reads passes
|
// New returns the server smallwebwaf runs: each request it reads passes
|
||||||
@@ -84,7 +93,7 @@ type Server struct {
|
|||||||
// applies the timeouts and size limits from then on.
|
// applies the timeouts and size limits from then on.
|
||||||
func New(params Params) *Server {
|
func New(params Params) *Server {
|
||||||
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
||||||
m := metrics.New(params.Config.MetricsTopN)
|
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName)
|
||||||
h := &handler{
|
h := &handler{
|
||||||
config: params.Config,
|
config: params.Config,
|
||||||
requestLog: params.RequestLog,
|
requestLog: params.RequestLog,
|
||||||
@@ -94,9 +103,12 @@ func New(params Params) *Server {
|
|||||||
now: params.Now,
|
now: params.Now,
|
||||||
metrics: m,
|
metrics: m,
|
||||||
limiter: ratelimit.New(ratelimit.Limits{
|
limiter: ratelimit.New(ratelimit.Limits{
|
||||||
PerMinute: params.Config.RateLimitPerMinute,
|
PerMinute: params.Config.RateLimitPerMinute,
|
||||||
PerHour: params.Config.RateLimitPerHour,
|
PerHour: params.Config.RateLimitPerHour,
|
||||||
PerDay: params.Config.RateLimitPerDay,
|
PerDay: params.Config.RateLimitPerDay,
|
||||||
|
BytesPerMinute: params.Config.BytesLimitPerMinute,
|
||||||
|
BytesPerHour: params.Config.BytesLimitPerHour,
|
||||||
|
BytesPerDay: params.Config.BytesLimitPerDay,
|
||||||
}),
|
}),
|
||||||
ledger: bans.New(bans.Rules{
|
ledger: bans.New(bans.Rules{
|
||||||
LimitBanDuration: params.Config.LimitBanDuration,
|
LimitBanDuration: params.Config.LimitBanDuration,
|
||||||
@@ -105,14 +117,24 @@ func New(params Params) *Server {
|
|||||||
AttackBanDuration: params.Config.AttackBanDuration,
|
AttackBanDuration: params.Config.AttackBanDuration,
|
||||||
MaxBans: params.Config.MaxBans,
|
MaxBans: params.Config.MaxBans,
|
||||||
}),
|
}),
|
||||||
geojs: lookup.New(lookup.Params{
|
lookupFile: params.LookupFile,
|
||||||
URL: params.GeoJSURL,
|
rules: params.Rules,
|
||||||
Now: params.Now,
|
alerts: params.Alerts,
|
||||||
ProcessLog: params.ProcessLog,
|
|
||||||
Metrics: m,
|
|
||||||
}),
|
|
||||||
rules: params.Rules,
|
|
||||||
}
|
}
|
||||||
|
h.geojs = lookup.New(lookup.Params{
|
||||||
|
URL: params.GeoJSURL,
|
||||||
|
Timeout: params.Config.LookupTimeout,
|
||||||
|
// The country lists, the headers and the biased thresholds act on
|
||||||
|
// the answer before the request goes on.
|
||||||
|
Wait: len(params.Config.DeniedCountries) > 0 ||
|
||||||
|
len(params.Config.ExclusivelyAllowedCountries) > 0 ||
|
||||||
|
params.Config.AddLookupHeaders || biasedThresholdsSet(params.Config),
|
||||||
|
Answered: h.addLookup,
|
||||||
|
Now: params.Now,
|
||||||
|
ProcessLog: params.ProcessLog,
|
||||||
|
Metrics: m,
|
||||||
|
Alerts: params.Alerts,
|
||||||
|
})
|
||||||
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
|
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
|
||||||
m.AddRules(params.Rules)
|
m.AddRules(params.Rules)
|
||||||
|
|
||||||
@@ -129,10 +151,11 @@ func New(params Params) *Server {
|
|||||||
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
||||||
ErrorLog: errorLog,
|
ErrorLog: errorLog,
|
||||||
},
|
},
|
||||||
Ledger: h.ledger,
|
Ledger: h.ledger,
|
||||||
Limiter: h.limiter,
|
Limiter: h.limiter,
|
||||||
GeoJS: h.geojs,
|
GeoJS: h.geojs,
|
||||||
Metrics: m,
|
LookupFile: h.lookupFile,
|
||||||
|
Metrics: m,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -149,7 +172,9 @@ type handler struct {
|
|||||||
limiter *ratelimit.Limiter
|
limiter *ratelimit.Limiter
|
||||||
ledger *bans.Ledger
|
ledger *bans.Ledger
|
||||||
geojs *lookup.GeoJS
|
geojs *lookup.GeoJS
|
||||||
|
lookupFile *lookup.File
|
||||||
rules *rules.Files
|
rules *rules.Files
|
||||||
|
alerts *alerts.Queue
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTransport returns what carries requests to the app. It never goes
|
// newTransport returns what carries requests to the app. It never goes
|
||||||
@@ -204,5 +229,10 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Once the response has ended, before the request is added to its
|
||||||
|
// client's history. Deferred, since ReverseProxy panics to end a
|
||||||
|
// response it cannot finish.
|
||||||
|
defer rq.countBytes()
|
||||||
|
|
||||||
rq.forward(r.Context())
|
rq.forward(r.Context())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,7 +14,9 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
@@ -65,6 +67,10 @@ const (
|
|||||||
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
||||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
||||||
|
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
||||||
|
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
||||||
|
lookupTimeout = "SWWAF_LOOKUP_TIMEOUT"
|
||||||
|
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
|
||||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||||
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
|
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
|
||||||
banResponse = "SWWAF_BAN_RESPONSE"
|
banResponse = "SWWAF_BAN_RESPONSE"
|
||||||
@@ -202,8 +208,8 @@ func startProxy(t *testing.T, appURL string, env map[string]string) (string, *ou
|
|||||||
return startProxyWithGeoJS(t, appURL, "", env)
|
return startProxyWithGeoJS(t, appURL, "", env)
|
||||||
}
|
}
|
||||||
|
|
||||||
// startProxyWithGeoJS is startProxy with clients' countries looked up at
|
// startProxyWithGeoJS is startProxy with clients' AS numbers and
|
||||||
// geojsURL.
|
// countries looked up at geojsURL.
|
||||||
func startProxyWithGeoJS(
|
func startProxyWithGeoJS(
|
||||||
t *testing.T, appURL, geojsURL string, env map[string]string,
|
t *testing.T, appURL, geojsURL string, env map[string]string,
|
||||||
) (string, *output) {
|
) (string, *output) {
|
||||||
@@ -216,43 +222,29 @@ func startProxyWithGeoJS(
|
|||||||
|
|
||||||
// startProxyWithClock is startProxyWithGeoJS with requests counted and
|
// startProxyWithClock is startProxyWithGeoJS with requests counted and
|
||||||
// bans made by the time now tells, and returns the server as well. Unless
|
// bans made by the time now tells, and returns the server as well. Unless
|
||||||
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
|
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules, and
|
||||||
|
// unless it sets SWWAF_INSTANCE_NAME, that is app, the label instance of
|
||||||
|
// every metric.
|
||||||
func startProxyWithClock(
|
func startProxyWithClock(
|
||||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||||
env map[string]string,
|
env map[string]string,
|
||||||
) (string, *output, *proxy.Server) {
|
) (string, *output, *proxy.Server) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()}
|
addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env)
|
||||||
maps.Copy(settings, env)
|
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
|
return addr, out, server
|
||||||
value, ok := settings[name]
|
}
|
||||||
|
|
||||||
return value, ok
|
// startProxyWithAlerts is startProxyWithClock, and returns the queue of
|
||||||
})
|
// the alerts the proxy raises as well, as newProxy makes them.
|
||||||
if err != nil {
|
func startProxyWithAlerts(
|
||||||
t.Fatalf("settings %v: %v", settings, err)
|
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||||
}
|
env map[string]string,
|
||||||
|
) (string, *output, *proxy.Server, *alerts.Queue) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
out := &output{}
|
server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env)
|
||||||
processLog := requestlog.NewProcessLogger(out)
|
|
||||||
|
|
||||||
ruleFiles, err := rules.Load(rules.Params{
|
|
||||||
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("rule files: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: out,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: geojsURL,
|
|
||||||
Now: now,
|
|
||||||
Rules: ruleFiles,
|
|
||||||
})
|
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -267,7 +259,83 @@ func startProxyWithClock(
|
|||||||
_ = server.Close()
|
_ = server.Close()
|
||||||
})
|
})
|
||||||
|
|
||||||
return listener.Addr().String(), out, server
|
return listener.Addr().String(), out, server, alertQueue
|
||||||
|
}
|
||||||
|
|
||||||
|
// newProxy makes the server startProxyWithClock starts, without starting
|
||||||
|
// it, and returns it, what it writes, and the queue of the alerts the
|
||||||
|
// proxy raises, as the settings in env make it. No alert is sent from the
|
||||||
|
// queue: they wait in it, for the test to look at. With no geojsURL, there
|
||||||
|
// is no stand-in for GeoJS to look clients up at, and SWWAF_LOOKUP_SOURCE
|
||||||
|
// is off unless env sets it. While it is file, the lookup database
|
||||||
|
// SWWAF_LOOKUP_DB_PATH names is read.
|
||||||
|
func newProxy(
|
||||||
|
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||||
|
env map[string]string,
|
||||||
|
) (*proxy.Server, *output, *alerts.Queue) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
settings := map[string]string{
|
||||||
|
"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir(), instanceName: "app",
|
||||||
|
}
|
||||||
|
if geojsURL == "" {
|
||||||
|
settings[lookupSource] = "off"
|
||||||
|
}
|
||||||
|
|
||||||
|
maps.Copy(settings, env)
|
||||||
|
|
||||||
|
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
|
||||||
|
value, ok := settings[name]
|
||||||
|
|
||||||
|
return value, ok
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("settings %v: %v", settings, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &output{}
|
||||||
|
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName)
|
||||||
|
|
||||||
|
ruleFiles, err := rules.Load(rules.Params{
|
||||||
|
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("rule files: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
alertQueue := alerts.New(alerts.Params{
|
||||||
|
WebhookURL: cfg.AlertWebhookURL,
|
||||||
|
Events: cfg.AlertEvents,
|
||||||
|
Cooldown: cfg.AlertCooldown,
|
||||||
|
MaxPerHour: cfg.AlertMaxPerHour,
|
||||||
|
Instance: cfg.InstanceName,
|
||||||
|
Now: now,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
|
||||||
|
var lookupFile *lookup.File
|
||||||
|
|
||||||
|
if cfg.LookupSource == fileSource {
|
||||||
|
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
||||||
|
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("lookup database: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
server := proxy.New(proxy.Params{
|
||||||
|
Config: cfg,
|
||||||
|
RequestLog: out,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
GeoJSURL: geojsURL,
|
||||||
|
LookupFile: lookupFile,
|
||||||
|
Now: now,
|
||||||
|
Rules: ruleFiles,
|
||||||
|
Alerts: alertQueue,
|
||||||
|
})
|
||||||
|
|
||||||
|
return server, out, alertQueue
|
||||||
}
|
}
|
||||||
|
|
||||||
// newClient returns an HTTP client that sends requests as they are made,
|
// newClient returns an HTTP client that sends requests as they are made,
|
||||||
|
|||||||
@@ -77,7 +77,8 @@ func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
|
|||||||
|
|
||||||
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
s, _, server := startWithClock(t, geojsURL, map[string]string{
|
||||||
rateLimitPerMinute: "1",
|
rateLimitPerMinute: "1",
|
||||||
rateLimitExemptPaths: "/assets/,/favicon.ico",
|
rateLimitExemptPaths: "/assets/,/favicon.ico",
|
||||||
denyNets: denied,
|
denyNets: denied,
|
||||||
|
|||||||
+91
-28
@@ -3,6 +3,7 @@ package proxy
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptrace"
|
"net/http/httptrace"
|
||||||
"net/http/httputil"
|
"net/http/httputil"
|
||||||
@@ -15,6 +16,7 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
@@ -48,7 +50,18 @@ type request struct {
|
|||||||
client netip.Addr
|
client netip.Addr
|
||||||
peer netip.Addr
|
peer netip.Addr
|
||||||
peerTrusted bool
|
peerTrusted bool
|
||||||
start time.Time
|
// lookedUp is true once the client's AS number and country have been
|
||||||
|
// looked up, whether or not an answer was there, and lookupAnswer is
|
||||||
|
// what the lookup gave then, the zero Answer while GeoJS had given none.
|
||||||
|
lookedUp bool
|
||||||
|
lookupAnswer lookup.Answer
|
||||||
|
// counted is true for a request the rate limits counted, whose bytes
|
||||||
|
// the byte limits count once it has ended. limitPercent and
|
||||||
|
// bytesPercent are then its client's limit percentages for the rate
|
||||||
|
// limits and for the byte limits.
|
||||||
|
counted bool
|
||||||
|
limitPercent, bytesPercent percentage
|
||||||
|
start time.Time
|
||||||
// checked is when the checks were done, and upstreamStart when the
|
// checked is when the checks were done, and upstreamStart when the
|
||||||
// request was handed to the app.
|
// request was handed to the app.
|
||||||
checked time.Time
|
checked time.Time
|
||||||
@@ -59,6 +72,9 @@ type request struct {
|
|||||||
refused atomic.Pointer[refusal]
|
refused atomic.Pointer[refusal]
|
||||||
// complete is true once the app's whole answer has been passed on.
|
// complete is true once the app's whole answer has been passed on.
|
||||||
complete bool
|
complete bool
|
||||||
|
// upgraded is the connection to the app once the app has switched
|
||||||
|
// protocols, as for a WebSocket, and nil otherwise.
|
||||||
|
upgraded *upgradedConn
|
||||||
|
|
||||||
// mu guards what follows. The timeouts run on goroutines of their
|
// mu guards what follows. The timeouts run on goroutines of their
|
||||||
// own, and the transport starts and stops them, and notes the times
|
// own, and the transport starts and stops them, and notes the times
|
||||||
@@ -192,14 +208,17 @@ func (rq *request) check(ctx context.Context) *refusal {
|
|||||||
|
|
||||||
// checkClient runs the checks on the request's client, and returns the
|
// checkClient runs the checks on the request's client, and returns the
|
||||||
// action of the first that refuses the request, or "" when none does. A
|
// action of the first that refuses the request, or "" when none does. A
|
||||||
// client in SWWAF_ALLOW_NETS skips them. For any other client,
|
// client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any
|
||||||
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
|
// other client, SWWAF_DENY_NETS comes first, then a ban on its netblock,
|
||||||
// client either refuses is not looked up, and then the country lists; a
|
// so that a client either refuses is not looked up, then the lookup of
|
||||||
// request any of them refuses is not counted for the rate limits. Then
|
// its AS number and country, and then the country lists; a request any of
|
||||||
// come the rate limits, unless the client is in
|
// them refuses is not counted for the rate limits. Then come the rate
|
||||||
// SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under
|
// limits, unless the client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the
|
||||||
// SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted,
|
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that
|
||||||
// and last the rule files. ctx is the request's own context.
|
// every other request is counted, each of them by the client's limit
|
||||||
|
// percentages, and last the rule files. A request exempt from the rate
|
||||||
|
// limits is exempt from the byte limits too. ctx is the request's own
|
||||||
|
// context.
|
||||||
func (rq *request) checkClient(ctx context.Context) string {
|
func (rq *request) checkClient(ctx context.Context) string {
|
||||||
cfg := rq.h.config
|
cfg := rq.h.config
|
||||||
if isInside(rq.client, cfg.AllowNets) {
|
if isInside(rq.client, cfg.AllowNets) {
|
||||||
@@ -216,13 +235,21 @@ func (rq *request) checkClient(ctx context.Context) string {
|
|||||||
return requestlog.ActionBanned
|
return requestlog.ActionBanned
|
||||||
}
|
}
|
||||||
|
|
||||||
if rq.countryDenied(ctx) {
|
rq.lookUp(ctx)
|
||||||
|
|
||||||
|
if rq.countryDenied() {
|
||||||
return requestlog.ActionCountryDenied
|
return requestlog.ActionCountryDenied
|
||||||
}
|
}
|
||||||
|
|
||||||
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
|
rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
|
||||||
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
!pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
||||||
if !exempt && rq.limitBroken(now) {
|
if rq.counted {
|
||||||
|
rq.limitPercent, rq.bytesPercent = limitPercentages(cfg, rq.line.ASN, rq.line.Country)
|
||||||
|
rq.line.LimitPercent, rq.line.LimitPercentSetting = rq.limitPercent.logged()
|
||||||
|
rq.line.BytesPercent, rq.line.BytesPercentSetting = rq.bytesPercent.logged()
|
||||||
|
}
|
||||||
|
|
||||||
|
if rq.counted && rq.limitBroken(now) {
|
||||||
return requestlog.ActionRateLimited
|
return requestlog.ActionRateLimited
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -287,7 +314,9 @@ func (rq *request) forward(ctx context.Context) {
|
|||||||
|
|
||||||
// rewrite makes the request the app receives: the client's request,
|
// rewrite makes the request the app receives: the client's request,
|
||||||
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
|
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
|
||||||
// the request's id set.
|
// the request's id set, without any X-Client-ASN or X-Client-Country the
|
||||||
|
// client sent, whatever SWWAF_ADD_LOOKUP_HEADERS says, and, while it is
|
||||||
|
// set, with the client's AS number and country in them.
|
||||||
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
||||||
upstream := rq.h.config.UpstreamURL
|
upstream := rq.h.config.UpstreamURL
|
||||||
pr.Out.URL.Scheme = upstream.Scheme
|
pr.Out.URL.Scheme = upstream.Scheme
|
||||||
@@ -297,6 +326,12 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
|||||||
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
||||||
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
||||||
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
|
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
|
||||||
|
pr.Out.Header.Del(asnHeader)
|
||||||
|
pr.Out.Header.Del(countryHeader)
|
||||||
|
|
||||||
|
if rq.h.config.AddLookupHeaders {
|
||||||
|
setLookupHeaders(pr.Out.Header, rq.line.ASN, rq.line.Country)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
||||||
@@ -307,11 +342,18 @@ func (rq *request) modifyResponse(res *http.Response) error {
|
|||||||
if res.StatusCode == http.StatusSwitchingProtocols {
|
if res.StatusCode == http.StatusSwitchingProtocols {
|
||||||
// An upgraded connection, such as a WebSocket, is not cut by the
|
// An upgraded connection, such as a WebSocket, is not cut by the
|
||||||
// timeouts. ReverseProxy writes this answer straight to the
|
// timeouts. ReverseProxy writes this answer straight to the
|
||||||
// connection it takes over, not through rq.out.
|
// connection it takes over, not through rq.out, and then copies
|
||||||
|
// what passes each way through res.Body, the connection to the app.
|
||||||
rq.stopTimers()
|
rq.stopTimers()
|
||||||
rq.out.status = res.StatusCode
|
rq.out.status = res.StatusCode
|
||||||
rq.line.Websocket = true
|
rq.line.Websocket = true
|
||||||
|
|
||||||
|
conn, ok := res.Body.(io.ReadWriteCloser)
|
||||||
|
if ok {
|
||||||
|
rq.upgraded = &upgradedConn{ReadWriteCloser: conn}
|
||||||
|
res.Body = rq.upgraded
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -414,10 +456,7 @@ func (rq *request) finish() {
|
|||||||
line.ResponseContentType = header.Get("Content-Type")
|
line.ResponseContentType = header.Get("Content-Type")
|
||||||
line.CacheControl = header.Get("Cache-Control")
|
line.CacheControl = header.Get("Cache-Control")
|
||||||
line.Location = header.Get("Location")
|
line.Location = header.Get("Location")
|
||||||
|
line.RequestBytes = rq.requestBytes()
|
||||||
if rq.body != nil {
|
|
||||||
line.RequestBytes = rq.body.bytes.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
// limit is the setting whose size or time limit the request passed.
|
// limit is the setting whose size or time limit the request passed.
|
||||||
var limit string
|
var limit string
|
||||||
@@ -474,24 +513,48 @@ func timing(start, end time.Time) *float64 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// addToHistory adds the request, which has ended, to its client's
|
// addToHistory adds the request, which has ended, to its client's
|
||||||
// history.
|
// history, and then, for a client that was looked up, the lookup
|
||||||
|
// database's answer about it, or the answer from GeoJS kept about it, to
|
||||||
|
// that history and to the notes of the bans on its netblock: an answer
|
||||||
|
// may have come before either was there, and one from GeoJS that comes
|
||||||
|
// later is added when it comes.
|
||||||
func (rq *request) addToHistory() {
|
func (rq *request) addToHistory() {
|
||||||
var requestBytes int64
|
|
||||||
if rq.body != nil {
|
|
||||||
requestBytes = rq.body.bytes.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
forwarded := !rq.upstreamStart.IsZero()
|
forwarded := !rq.upstreamStart.IsZero()
|
||||||
|
group := clientGroup(rq.client)
|
||||||
|
|
||||||
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
|
rq.h.limiter.AddToHistory(group, rq.h.now(), ratelimit.Request{
|
||||||
Country: rq.line.Country,
|
|
||||||
Forwarded: forwarded,
|
Forwarded: forwarded,
|
||||||
Refused: !forwarded && rq.refused.Load() != nil,
|
Refused: !forwarded && rq.refused.Load() != nil,
|
||||||
Status: rq.out.status,
|
Status: rq.out.status,
|
||||||
RequestBytes: requestBytes,
|
RequestBytes: rq.requestBytes(),
|
||||||
ResponseBytes: rq.out.bytes,
|
ResponseBytes: rq.out.bytes,
|
||||||
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
if !rq.lookedUp {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// The lookup database's answer was there at once.
|
||||||
|
if rq.h.config.LookupSource == "file" {
|
||||||
|
rq.h.addLookup(rq.lookupAnswer)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
answer, kept := rq.h.geojs.Kept(group)
|
||||||
|
if kept {
|
||||||
|
rq.h.addLookup(answer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// requestBytes is how many bytes of the request's body have been read.
|
||||||
|
func (rq *request) requestBytes() int64 {
|
||||||
|
if rq.body == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
return rq.body.bytes.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
// clientRequestDeadline is when the client must have sent its whole
|
// clientRequestDeadline is when the client must have sent its whole
|
||||||
|
|||||||
@@ -119,7 +119,10 @@ func wantFullLine(t *testing.T, line logLine) {
|
|||||||
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
|
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
|
||||||
CacheControl: "no-store", Location: "/elsewhere",
|
CacheControl: "no-store", Location: "/elsewhere",
|
||||||
Action: requestlog.ActionForward,
|
Action: requestlog.ActionForward,
|
||||||
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
// Its 3 bytes in and 5 out, each way counted by default.
|
||||||
|
Counts: ratelimit.Counts{
|
||||||
|
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
|
||||||
|
},
|
||||||
})
|
})
|
||||||
if !reflect.DeepEqual(line.Line, want) {
|
if !reflect.DeepEqual(line.Line, want) {
|
||||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||||
|
|||||||
@@ -197,15 +197,15 @@ func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
|
|||||||
|
|
||||||
metrics := s.scrape(scraper)
|
metrics := s.scrape(scraper)
|
||||||
wantMetric(t, metrics,
|
wantMetric(t, metrics,
|
||||||
`smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
|
`smallwebwaf_rule_matches_total{action="block",instance="app",rule_id="blocked"}`, 1)
|
||||||
wantMetric(t, metrics,
|
wantMetric(t, metrics,
|
||||||
`smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
|
`smallwebwaf_rule_matches_total{action="ban",instance="app",rule_id="probe"}`, 1)
|
||||||
wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
|
wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2)
|
||||||
wantMetric(t, metrics,
|
wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+
|
||||||
`smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
|
`instance="app",status_class="4xx"}`, 1)
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1)
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0)
|
||||||
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
|
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
// writeRules writes content as a rule file into a new directory, and
|
// writeRules writes content as a rule file into a new directory, and
|
||||||
|
|||||||
@@ -10,8 +10,9 @@ import (
|
|||||||
// checkRules checks the request against the rules of the rule files at
|
// checkRules checks the request against the rules of the rule files at
|
||||||
// now, notes the ids of those it matches in the log line, and returns the
|
// now, notes the ids of those it matches in the log line, and returns the
|
||||||
// action of the rule that refuses it, ActionRuleBlocked for a block rule
|
// action of the rule that refuses it, ActionRuleBlocked for a block rule
|
||||||
// and ActionBanned for a ban rule, or "" when none does. In enforce mode
|
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans
|
||||||
// a ban rule bans the client's netblock for a clear sign of attack.
|
// the client's netblock for a clear sign of attack, or in observe mode
|
||||||
|
// raises the alert for the ban it would have made.
|
||||||
func (rq *request) checkRules(now time.Time) string {
|
func (rq *request) checkRules(now time.Time) string {
|
||||||
matched := rq.h.rules.Match(rq.in)
|
matched := rq.h.rules.Match(rq.in)
|
||||||
|
|
||||||
@@ -29,9 +30,7 @@ func (rq *request) checkRules(now time.Time) string {
|
|||||||
case rules.ActionBlock:
|
case rules.ActionBlock:
|
||||||
return requestlog.ActionRuleBlocked
|
return requestlog.ActionRuleBlocked
|
||||||
case rules.ActionBan:
|
case rules.ActionBan:
|
||||||
if !rq.h.config.Observe {
|
rq.banForAttack(now, last)
|
||||||
rq.banForAttack(now, last)
|
|
||||||
}
|
|
||||||
|
|
||||||
return requestlog.ActionBanned
|
return requestlog.ActionBanned
|
||||||
default:
|
default:
|
||||||
|
|||||||
@@ -16,10 +16,10 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
|
|||||||
start := midnight()
|
start := midnight()
|
||||||
|
|
||||||
for i, r := range []ratelimit.Request{
|
for i, r := range []ratelimit.Request{
|
||||||
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
|
{Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
|
||||||
{Forwarded: true, Status: 101},
|
{Forwarded: true, Status: 101},
|
||||||
{Forwarded: true, Status: 304, RequestBytes: 5},
|
{Forwarded: true, Status: 304, RequestBytes: 5},
|
||||||
{Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
|
{Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
|
||||||
{Forwarded: true, Status: 502, ResponseBytes: 12},
|
{Forwarded: true, Status: 502, ResponseBytes: 12},
|
||||||
// Closed without an answer: refused, and no response.
|
// Closed without an answer: refused, and no response.
|
||||||
{Refused: true, Status: 0},
|
{Refused: true, Status: 0},
|
||||||
@@ -33,8 +33,6 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
|
|||||||
want := ratelimit.History{
|
want := ratelimit.History{
|
||||||
FirstSeen: start,
|
FirstSeen: start,
|
||||||
LastSeen: start.Add(6 * time.Minute),
|
LastSeen: start.Add(6 * time.Minute),
|
||||||
Country: "FR",
|
|
||||||
LookedUp: start.Add(3 * time.Minute),
|
|
||||||
Requests: 7,
|
Requests: 7,
|
||||||
Forwarded: 4,
|
Forwarded: 4,
|
||||||
Refused: 2,
|
Refused: 2,
|
||||||
@@ -52,6 +50,43 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
other := netip.MustParsePrefix("198.51.100.7/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
||||||
|
limiter.AddLookup(client, start, "AS64496", "Example Net", "DE")
|
||||||
|
|
||||||
|
// A later answer replaces it, and one for a client the table does not
|
||||||
|
// hold adds no client.
|
||||||
|
limiter.AddLookup(client, start.Add(time.Hour), "AS64497", "Other Net", "FR")
|
||||||
|
limiter.AddLookup(other, start, "AS64496", "Example Net", "DE")
|
||||||
|
|
||||||
|
want := ratelimit.History{
|
||||||
|
FirstSeen: start,
|
||||||
|
LastSeen: start,
|
||||||
|
ASN: "AS64497",
|
||||||
|
ASName: "Other Net",
|
||||||
|
Country: "FR",
|
||||||
|
LookedUp: start.Add(time.Hour),
|
||||||
|
Requests: 1,
|
||||||
|
Forwarded: 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
got := historyOf(t, limiter, client)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("history\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if clients := limiter.Snapshot(); len(clients) != 1 {
|
||||||
|
t.Errorf("the table holds %+v, want %s alone", clients, client)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestResetKeepsTheHistory(t *testing.T) {
|
func TestResetKeepsTheHistory(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
|||||||
+200
-89
@@ -1,9 +1,9 @@
|
|||||||
// Package ratelimit keeps the table of clients: each client's requests
|
// Package ratelimit keeps the table of clients: each client's requests
|
||||||
// counted over a minute, an hour and a day, as the "Counting method"
|
// and bytes counted over a minute, an hour and a day, as the "Counting
|
||||||
// section of SPEC.md describes, which tell when a request takes the client
|
// method" section of SPEC.md describes, which tell when a request takes
|
||||||
// over a rate limit, and each client's history since it was first seen.
|
// the client over a rate limit or a byte limit, and each client's history
|
||||||
// At most 20,000 clients are kept, in memory, and written to clients.json
|
// since it was first seen. At most 20,000 clients are kept, in memory, and
|
||||||
// and read from it by the state package.
|
// written to clients.json and read from it by the state package.
|
||||||
package ratelimit
|
package ratelimit
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -23,19 +23,30 @@ const maxClients = 20000
|
|||||||
|
|
||||||
const day = 24 * time.Hour
|
const day = 24 * time.Hour
|
||||||
|
|
||||||
|
// The kinds of limits, as the metrics name them.
|
||||||
|
const (
|
||||||
|
// KindRequests is a rate limit, on a client's requests.
|
||||||
|
KindRequests = "requests"
|
||||||
|
// KindBytes is a byte limit, on a client's bytes.
|
||||||
|
KindBytes = "bytes"
|
||||||
|
)
|
||||||
|
|
||||||
// Limits are the most requests a client may make in a minute, an hour and
|
// Limits are the most requests a client may make in a minute, an hour and
|
||||||
// a day. Zero is no limit.
|
// a day, and the most bytes. Zero is no limit.
|
||||||
type Limits struct {
|
type Limits struct {
|
||||||
PerMinute int64
|
PerMinute int64
|
||||||
PerHour int64
|
PerHour int64
|
||||||
PerDay int64
|
PerDay int64
|
||||||
|
BytesPerMinute int64
|
||||||
|
BytesPerHour int64
|
||||||
|
BytesPerDay int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// Limiter counts each client's requests against the limits, and keeps
|
// Limiter counts each client's requests and bytes against the limits, and
|
||||||
// its history. It is safe for concurrent use.
|
// keeps its history. It is safe for concurrent use.
|
||||||
type Limiter struct {
|
type Limiter struct {
|
||||||
// windows are the minute, the hour and the day, in the order of
|
// windows are the minute, the hour and the day, in the order of
|
||||||
// Client.buckets.
|
// Client.buckets and Client.byteBuckets.
|
||||||
windows [3]window
|
windows [3]window
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
@@ -43,17 +54,23 @@ type Limiter struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Client is a client in the table, as clients.json holds it: its buckets
|
// Client is a client in the table, as clients.json holds it: its buckets
|
||||||
// in each window, and its history.
|
// of requests and of bytes in each window, and its history.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
type Client struct {
|
type Client struct {
|
||||||
Client netip.Prefix `json:"client"`
|
Client netip.Prefix `json:"client"`
|
||||||
Minute Buckets `json:"minute"`
|
Minute Buckets `json:"minute"`
|
||||||
Hour Buckets `json:"hour"`
|
Hour Buckets `json:"hour"`
|
||||||
Day Buckets `json:"day"`
|
Day Buckets `json:"day"`
|
||||||
History History `json:"history"`
|
MinuteBytes Buckets `json:"minute_bytes"`
|
||||||
|
HourBytes Buckets `json:"hour_bytes"`
|
||||||
|
DayBytes Buckets `json:"day_bytes"`
|
||||||
|
History History `json:"history"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Buckets are a client's two buckets in one window: the requests in the
|
// Buckets are a client's two buckets in one window: the requests, or the
|
||||||
// bucket under way, which began at Start, and in the bucket before it.
|
// bytes, in the bucket under way, which began at Start, and in the bucket
|
||||||
|
// before it.
|
||||||
type Buckets struct {
|
type Buckets struct {
|
||||||
Start time.Time `json:"start"`
|
Start time.Time `json:"start"`
|
||||||
Current int64 `json:"current"`
|
Current int64 `json:"current"`
|
||||||
@@ -66,8 +83,12 @@ type Buckets struct {
|
|||||||
type History struct {
|
type History struct {
|
||||||
FirstSeen time.Time `json:"first_seen"`
|
FirstSeen time.Time `json:"first_seen"`
|
||||||
LastSeen time.Time `json:"last_seen"`
|
LastSeen time.Time `json:"last_seen"`
|
||||||
// Country is the client's country as it was last looked up, and
|
// ASN, ASName and Country are the client's AS number, AS name and
|
||||||
// LookedUp when that was; both are empty while it never was.
|
// country as last looked up, each empty when the lookup could not
|
||||||
|
// find it, and LookedUp is when the lookup gave that answer; all are
|
||||||
|
// empty while the client never was looked up.
|
||||||
|
ASN string `json:"asn,omitempty"`
|
||||||
|
ASName string `json:"as_name,omitempty"`
|
||||||
Country string `json:"country,omitempty"`
|
Country string `json:"country,omitempty"`
|
||||||
LookedUp time.Time `json:"looked_up,omitzero"`
|
LookedUp time.Time `json:"looked_up,omitzero"`
|
||||||
// Requests are all the client's requests: Forwarded those passed to
|
// Requests are all the client's requests: Forwarded those passed to
|
||||||
@@ -97,14 +118,12 @@ type Responses struct {
|
|||||||
|
|
||||||
// Offences are a client's offences, by kind.
|
// Offences are a client's offences, by kind.
|
||||||
type Offences struct {
|
type Offences struct {
|
||||||
// Limit is its requests that broke a rate limit.
|
// Limit is its requests that broke a rate limit or a byte limit.
|
||||||
Limit int64 `json:"limit"`
|
Limit int64 `json:"limit"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Request is what a client's history keeps of one of its requests.
|
// Request is what a client's history keeps of one of its requests.
|
||||||
type Request struct {
|
type Request struct {
|
||||||
// Country is the client's country, when the request looked it up.
|
|
||||||
Country string
|
|
||||||
// Forwarded is true for a request passed to the app, Refused for one
|
// Forwarded is true for a request passed to the app, Refused for one
|
||||||
// refused before anything reached it, a 401 at smallwebwaf's own
|
// refused before anything reached it, a 401 at smallwebwaf's own
|
||||||
// endpoints included. Both are false for any other request smallwebwaf
|
// endpoints included. Both are false for any other request smallwebwaf
|
||||||
@@ -117,7 +136,8 @@ type Request struct {
|
|||||||
// and of its response.
|
// and of its response.
|
||||||
RequestBytes int64
|
RequestBytes int64
|
||||||
ResponseBytes int64
|
ResponseBytes int64
|
||||||
// BrokeLimit is true for a request that broke a rate limit.
|
// BrokeLimit is true for a request that broke a rate limit or a byte
|
||||||
|
// limit.
|
||||||
BrokeLimit bool
|
BrokeLimit bool
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -130,62 +150,74 @@ func New(limits Limits) *Limiter {
|
|||||||
|
|
||||||
return &Limiter{
|
return &Limiter{
|
||||||
windows: [3]window{
|
windows: [3]window{
|
||||||
{name: "minute", length: time.Minute, limit: limits.PerMinute},
|
{
|
||||||
{name: "hour", length: time.Hour, limit: limits.PerHour},
|
name: "minute", length: time.Minute,
|
||||||
{name: "day", length: day, limit: limits.PerDay},
|
limit: limits.PerMinute, byteLimit: limits.BytesPerMinute,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "hour", length: time.Hour,
|
||||||
|
limit: limits.PerHour, byteLimit: limits.BytesPerHour,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "day", length: day,
|
||||||
|
limit: limits.PerDay, byteLimit: limits.BytesPerDay,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
clients: clients,
|
clients: clients,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Hit is a request that takes a client over a rate limit.
|
// Hit is a request that takes a client over a rate limit, or whose bytes
|
||||||
|
// take it over a byte limit.
|
||||||
type Hit struct {
|
type Hit struct {
|
||||||
|
// Kind is KindRequests for a rate limit, KindBytes for a byte limit.
|
||||||
|
Kind string
|
||||||
// Window is "minute", "hour" or "day".
|
// Window is "minute", "hour" or "day".
|
||||||
Window string
|
Window string
|
||||||
// Limit is the window's limit.
|
// Limit is the window's limit, as the client's percentage of it.
|
||||||
Limit int64
|
Limit int64
|
||||||
// Requests is the client's requests counted in the window, this one
|
// Count is the client's requests, or bytes, counted in the window,
|
||||||
// included.
|
// this request's included.
|
||||||
Requests float64
|
Count float64
|
||||||
}
|
}
|
||||||
|
|
||||||
// Counts are a client's requests in the minute, the hour and the day that
|
// Counts are a client's requests and bytes in the minute, the hour and
|
||||||
// end at a request, that request included.
|
// the day that end at a request, that request's included.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
||||||
type Counts struct {
|
type Counts struct {
|
||||||
Minute float64 `json:"minute"`
|
Minute float64 `json:"minute"`
|
||||||
Hour float64 `json:"hour"`
|
Hour float64 `json:"hour"`
|
||||||
Day float64 `json:"day"`
|
Day float64 `json:"day"`
|
||||||
|
MinuteBytes float64 `json:"minute_bytes"`
|
||||||
|
HourBytes float64 `json:"hour_bytes"`
|
||||||
|
DayBytes float64 `json:"day_bytes"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Count counts a request from client at now, in every window, whether or
|
// Count counts a request from client at now, in every window, whether or
|
||||||
// not it is refused, and returns the client's requests in each window. It
|
// not it is refused, and returns the client's counts in each window. It
|
||||||
// reports whether the request takes the client over a limit, and the
|
// reports whether the request takes the client over a rate limit, of
|
||||||
// window whose limit it goes over, the shortest if it is over several.
|
// which the client gets the percentage percent, rounded down, and the hit:
|
||||||
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
|
// the window whose limit it goes over, the shortest if it is over
|
||||||
l.mu.Lock()
|
// several. A limit that is off stays off.
|
||||||
defer l.mu.Unlock()
|
func (l *Limiter) Count(
|
||||||
|
client netip.Prefix, now time.Time, percent int64,
|
||||||
var (
|
) (Counts, Hit, bool) {
|
||||||
requests [3]float64
|
return l.count(client, now, 1, 0, percent)
|
||||||
hit Hit
|
|
||||||
)
|
|
||||||
|
|
||||||
for i, b := range l.get(client).buckets() {
|
|
||||||
w := l.windows[i]
|
|
||||||
|
|
||||||
requests[i] = b.add(now, w.length)
|
|
||||||
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
|
|
||||||
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
|
|
||||||
|
|
||||||
return counts, hit, hit.Window != ""
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reset sets client's counts in every window back to zero. Its history
|
// CountBytes counts bytes, those of a request from client that has ended,
|
||||||
// keeps its totals.
|
// at now, in every window, and returns the client's counts in each window.
|
||||||
|
// It reports whether the bytes take the client over a byte limit, of which
|
||||||
|
// the client gets the percentage percent, and the hit, as Count does.
|
||||||
|
func (l *Limiter) CountBytes(
|
||||||
|
client netip.Prefix, now time.Time, bytes, percent int64,
|
||||||
|
) (Counts, Hit, bool) {
|
||||||
|
return l.count(client, now, 0, bytes, percent)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset sets client's counts of requests and of bytes in every window
|
||||||
|
// back to zero. Its history keeps its totals.
|
||||||
func (l *Limiter) Reset(client netip.Prefix) {
|
func (l *Limiter) Reset(client netip.Prefix) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
@@ -193,6 +225,7 @@ func (l *Limiter) Reset(client netip.Prefix) {
|
|||||||
c, seen := l.clients.Peek(client)
|
c, seen := l.clients.Peek(client)
|
||||||
if seen {
|
if seen {
|
||||||
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
|
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
|
||||||
|
c.MinuteBytes, c.HourBytes, c.DayBytes = Buckets{}, Buckets{}, Buckets{}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -209,11 +242,6 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
|
|||||||
|
|
||||||
h.LastSeen = now
|
h.LastSeen = now
|
||||||
|
|
||||||
if r.Country != "" {
|
|
||||||
h.Country = r.Country
|
|
||||||
h.LookedUp = now
|
|
||||||
}
|
|
||||||
|
|
||||||
h.Requests++
|
h.Requests++
|
||||||
if r.Forwarded {
|
if r.Forwarded {
|
||||||
h.Forwarded++
|
h.Forwarded++
|
||||||
@@ -232,6 +260,25 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AddLookup gives client's history its AS number, AS name and country, as
|
||||||
|
// the lookup gave them at lookedUp, if the table of clients holds the
|
||||||
|
// client.
|
||||||
|
// It does not make the client the most recently seen.
|
||||||
|
func (l *Limiter) AddLookup(
|
||||||
|
client netip.Prefix, lookedUp time.Time, asn, asName, country string,
|
||||||
|
) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
c, held := l.clients.Peek(client)
|
||||||
|
if !held {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
h := &c.History
|
||||||
|
h.ASN, h.ASName, h.Country, h.LookedUp = asn, asName, country, lookedUp
|
||||||
|
}
|
||||||
|
|
||||||
// Requests returns how many requests the clients inside netblock have
|
// Requests returns how many requests the clients inside netblock have
|
||||||
// sent, as their histories count them.
|
// sent, as their histories count them.
|
||||||
func (l *Limiter) Requests(netblock netip.Prefix) int64 {
|
func (l *Limiter) Requests(netblock netip.Prefix) int64 {
|
||||||
@@ -312,12 +359,13 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
|
|||||||
l.clients.Purge()
|
l.clients.Purge()
|
||||||
|
|
||||||
for _, c := range clients {
|
for _, c := range clients {
|
||||||
for i, b := range c.buckets() {
|
for i, w := range l.windows {
|
||||||
// The window that ends at now covers neither bucket once it
|
for _, b := range []*Buckets{c.buckets()[i], c.byteBuckets()[i]} {
|
||||||
// begins after the bucket under way has ended.
|
// The window that ends at now covers neither bucket once it
|
||||||
length := l.windows[i].length
|
// begins after the bucket under way has ended.
|
||||||
if !now.Add(-length).Before(b.Start.Add(length)) {
|
if !now.Add(-w.length).Before(b.Start.Add(w.length)) {
|
||||||
*b = Buckets{}
|
*b = Buckets{}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -325,6 +373,51 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// count adds requests and bytes from client at now to its buckets in
|
||||||
|
// every window, and returns its counts. A limit is broken only by what is
|
||||||
|
// added to it, so that a request whose bytes are counted after another of
|
||||||
|
// the client's requests broke a rate limit does not break it too. The
|
||||||
|
// client gets the percentage percent of each limit.
|
||||||
|
func (l *Limiter) count(
|
||||||
|
client netip.Prefix, now time.Time, requests, bytes, percent int64,
|
||||||
|
) (Counts, Hit, bool) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
c := l.get(client)
|
||||||
|
requestBuckets, byteBuckets := c.buckets(), c.byteBuckets()
|
||||||
|
|
||||||
|
var (
|
||||||
|
requestCounts, byteCounts [3]float64
|
||||||
|
hit Hit
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, w := range l.windows {
|
||||||
|
requestCounts[i] = requestBuckets[i].add(now, w.length, requests)
|
||||||
|
byteCounts[i] = byteBuckets[i].add(now, w.length, bytes)
|
||||||
|
limit, byteLimit := percentOf(w.limit, percent), percentOf(w.byteLimit, percent)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case hit.Window != "":
|
||||||
|
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(limit):
|
||||||
|
hit = Hit{
|
||||||
|
Kind: KindRequests, Window: w.name, Limit: limit, Count: requestCounts[i],
|
||||||
|
}
|
||||||
|
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(byteLimit):
|
||||||
|
hit = Hit{
|
||||||
|
Kind: KindBytes, Window: w.name, Limit: byteLimit, Count: byteCounts[i],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
counts := Counts{
|
||||||
|
Minute: requestCounts[0], Hour: requestCounts[1], Day: requestCounts[2],
|
||||||
|
MinuteBytes: byteCounts[0], HourBytes: byteCounts[1], DayBytes: byteCounts[2],
|
||||||
|
}
|
||||||
|
|
||||||
|
return counts, hit, hit.Window != ""
|
||||||
|
}
|
||||||
|
|
||||||
// get returns client's entry in the table, a new one if it has none, and
|
// get returns client's entry in the table, a new one if it has none, and
|
||||||
// makes it the most recently seen.
|
// makes it the most recently seen.
|
||||||
func (l *Limiter) get(client netip.Prefix) *Client {
|
func (l *Limiter) get(client netip.Prefix) *Client {
|
||||||
@@ -337,30 +430,48 @@ func (l *Limiter) get(client netip.Prefix) *Client {
|
|||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
// buckets returns c's buckets in the minute, the hour and the day.
|
// buckets returns c's buckets of requests in the minute, the hour and the
|
||||||
|
// day.
|
||||||
func (c *Client) buckets() [3]*Buckets {
|
func (c *Client) buckets() [3]*Buckets {
|
||||||
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
|
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
|
||||||
}
|
}
|
||||||
|
|
||||||
// window is a length of time over which requests are counted, and the
|
// byteBuckets returns c's buckets of bytes in the minute, the hour and the
|
||||||
// most requests a client may make in it.
|
// day.
|
||||||
type window struct {
|
func (c *Client) byteBuckets() [3]*Buckets {
|
||||||
name string
|
return [3]*Buckets{&c.MinuteBytes, &c.HourBytes, &c.DayBytes}
|
||||||
length time.Duration
|
|
||||||
limit int64
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// add counts a request at now in a window of length, and returns the
|
// window is a length of time over which requests and bytes are counted,
|
||||||
// client's requests in the window that ends at now: those in the bucket
|
// and the most requests and the most bytes a client may have in it.
|
||||||
// under way, and those in the bucket before it weighted by how much of
|
type window struct {
|
||||||
// that bucket the window still covers.
|
name string
|
||||||
|
length time.Duration
|
||||||
|
limit int64
|
||||||
|
byteLimit int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// percentOf returns the percentage percent of limit, rounded down. It is
|
||||||
|
// written as limit's hundreds times percent, plus the rest's share, since
|
||||||
|
// limit*percent can overflow for a byte limit.
|
||||||
|
func percentOf(limit, percent int64) int64 {
|
||||||
|
const hundred = 100
|
||||||
|
|
||||||
|
return limit/hundred*percent + limit%hundred*percent/hundred
|
||||||
|
}
|
||||||
|
|
||||||
|
// add counts n requests, or n bytes, at now in a window of length, and
|
||||||
|
// returns the client's count in the window that ends at now: what is in
|
||||||
|
// the bucket under way, and what is in the bucket before it weighted by
|
||||||
|
// how much of that bucket the window still covers. With n zero it counts
|
||||||
|
// nothing, and returns the count.
|
||||||
//
|
//
|
||||||
// Concurrent requests can be counted out of order, so now can be a moment
|
// Concurrent requests can be counted out of order, so now can be a moment
|
||||||
// before the bucket under way began; such a request is counted in that
|
// before the bucket under way began; such a request is counted in that
|
||||||
// bucket. A request dated more than a second before it means the clock
|
// bucket. A request dated more than a second before it means the clock
|
||||||
// was set back, and the buckets start afresh: otherwise the bucket before
|
// was set back, and the buckets start afresh: otherwise the bucket before
|
||||||
// would keep its full weight until the clock caught up.
|
// would keep its full weight until the clock caught up.
|
||||||
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
func (b *Buckets) add(now time.Time, length time.Duration, n int64) float64 {
|
||||||
if now.Before(b.Start.Add(-time.Second)) {
|
if now.Before(b.Start.Add(-time.Second)) {
|
||||||
*b = Buckets{}
|
*b = Buckets{}
|
||||||
}
|
}
|
||||||
@@ -377,7 +488,7 @@ func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
|||||||
b.Current = 0
|
b.Current = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.Current++
|
b.Current += n
|
||||||
|
|
||||||
elapsed := max(now.Sub(b.Start), 0)
|
elapsed := max(now.Sub(b.Start), 0)
|
||||||
covered := 1 - float64(elapsed)/float64(length)
|
covered := 1 - float64(elapsed)/float64(length)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package ratelimit_test
|
package ratelimit_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"math"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@@ -11,6 +12,10 @@ import (
|
|||||||
// limit is the limit the tests set.
|
// limit is the limit the tests set.
|
||||||
const limit = 3
|
const limit = 3
|
||||||
|
|
||||||
|
// whole is the percentage of each limit a client gets when nothing lowers
|
||||||
|
// its limits.
|
||||||
|
const whole = 100
|
||||||
|
|
||||||
// The windows, as Count names them.
|
// The windows, as Count names them.
|
||||||
const (
|
const (
|
||||||
minute = "minute"
|
minute = "minute"
|
||||||
@@ -62,22 +67,180 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
|
|||||||
start := midnight()
|
start := midnight()
|
||||||
|
|
||||||
for range limit {
|
for range limit {
|
||||||
_, _, over := limiter.Count(client, start)
|
_, _, over := limiter.Count(client, start, whole)
|
||||||
if over {
|
if over {
|
||||||
t.Fatal("a request within the limit is over it")
|
t.Fatal("a request within the limit is over it")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Over both limits; the minute's is named, with the four requests.
|
// Over both limits; the minute's is named, with the four requests.
|
||||||
_, hit, over := limiter.Count(client, start)
|
_, hit, over := limiter.Count(client, start, whole)
|
||||||
|
|
||||||
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
|
want := ratelimit.Hit{
|
||||||
|
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
|
||||||
|
}
|
||||||
if !over || hit != want {
|
if !over || hit != want {
|
||||||
t.Errorf("request over the limit gives %+v and %t, want %+v and true",
|
t.Errorf("request over the limit gives %+v and %t, want %+v and true",
|
||||||
hit, over, want)
|
hit, over, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestClientGetsItsPercentageOfEachLimitRoundedDown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 5, BytesPerDay: math.MaxInt64})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
// Half of 5 requests is 2.5, rounded down to 2: the third is over.
|
||||||
|
for range 2 {
|
||||||
|
_, _, over := limiter.Count(client, start, 50)
|
||||||
|
if over {
|
||||||
|
t.Fatal("a request within half the limit is over it")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_, hit, over := limiter.Count(client, start, 50)
|
||||||
|
|
||||||
|
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: minute, Limit: 2, Count: 3}
|
||||||
|
if !over || hit != want {
|
||||||
|
t.Errorf("the third request gives %+v and %t, want %+v and true", hit, over, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Half of the largest byte limit is still far above a TiB: working it
|
||||||
|
// out does not overflow.
|
||||||
|
_, hit, over = limiter.CountBytes(client, start, 1<<40, 50)
|
||||||
|
if over {
|
||||||
|
t.Errorf("a TiB is over half the largest byte limit: %+v", hit)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestZeroPercentIsAZeroAllowanceAndALimitOffStaysOff(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Only the hour has limits: the minute's and the day's are off.
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit, BytesPerHour: 1000})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
// At 0 percent, the first request and the first byte are over the
|
||||||
|
// hour's limits, which are 0; the minute's, which are off, stay off.
|
||||||
|
_, hit, _ := limiter.Count(client, start, 0)
|
||||||
|
|
||||||
|
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: hour, Limit: 0, Count: 1}
|
||||||
|
if hit != want {
|
||||||
|
t.Errorf("the first request gives %+v, want %+v", hit, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, hit, _ = limiter.CountBytes(client, start, 1, 0)
|
||||||
|
|
||||||
|
want = ratelimit.Hit{Kind: ratelimit.KindBytes, Window: hour, Limit: 0, Count: 1}
|
||||||
|
if hit != want {
|
||||||
|
t.Errorf("the first byte gives %+v, want %+v", hit, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEachByteLimitIsBrokenByTheBytesCounted(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const byteLimit = 1000
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
window string
|
||||||
|
limits ratelimit.Limits
|
||||||
|
}{
|
||||||
|
{minute, ratelimit.Limits{BytesPerMinute: byteLimit}},
|
||||||
|
{hour, ratelimit.Limits{BytesPerHour: byteLimit}},
|
||||||
|
{"day", ratelimit.Limits{BytesPerDay: byteLimit}},
|
||||||
|
} {
|
||||||
|
t.Run(tc.window, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(tc.limits)
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
// 600 bytes are within the limit, 600 more over it.
|
||||||
|
_, _, over := limiter.CountBytes(client, midnight(), 600, whole)
|
||||||
|
if over {
|
||||||
|
t.Fatal("600 bytes are over the limit of 1000")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, hit, over := limiter.CountBytes(client, midnight(), 600, whole)
|
||||||
|
|
||||||
|
want := ratelimit.Hit{
|
||||||
|
Kind: ratelimit.KindBytes, Window: tc.window, Limit: byteLimit, Count: 1200,
|
||||||
|
}
|
||||||
|
if !over || hit != want {
|
||||||
|
t.Errorf("1200 bytes give %+v and %t, want %+v and true", hit, over, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestALimitIsBrokenOnlyByWhatIsAddedToIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 2, BytesPerMinute: 1000})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
other := netip.MustParsePrefix("203.0.113.10/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
// The third request breaks the rate limit. The bytes of a request
|
||||||
|
// counted after it, within the byte limit, do not break it again.
|
||||||
|
for range 2 {
|
||||||
|
wantCount(t, limiter, client, start, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
wantCount(t, limiter, client, start, minute)
|
||||||
|
wantBytesCount(t, limiter, client, start, 500, "")
|
||||||
|
wantBytesCount(t, limiter, client, start, 600, ratelimit.KindBytes)
|
||||||
|
|
||||||
|
// Bytes over the byte limit do not have the next request break it, nor
|
||||||
|
// the rate limit, which that request is within.
|
||||||
|
wantBytesCount(t, limiter, other, start, 1200, ratelimit.KindBytes)
|
||||||
|
wantCount(t, limiter, other, start, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCountGivesTheBytesInEachWindow(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
limiter.CountBytes(client, start, 300, whole)
|
||||||
|
|
||||||
|
// A quarter into the next hour, the minute has only these 100 bytes.
|
||||||
|
// The hour still covers three quarters of the bucket before, whose 300
|
||||||
|
// bytes count 225, and these: 325. The day covers all 400.
|
||||||
|
later := start.Add(time.Hour + time.Hour/4)
|
||||||
|
limiter.CountBytes(client, later, 100, whole)
|
||||||
|
|
||||||
|
// A request's counts give the bytes counted so far too.
|
||||||
|
counts, _, _ := limiter.Count(client, later, whole)
|
||||||
|
|
||||||
|
want := ratelimit.Counts{
|
||||||
|
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 100, HourBytes: 325, DayBytes: 400,
|
||||||
|
}
|
||||||
|
if counts != want {
|
||||||
|
t.Errorf("counts %+v, want %+v", counts, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResetSetsTheBytesBackToZero(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{BytesPerDay: 1000})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
wantBytesCount(t, limiter, client, start, 1200, ratelimit.KindBytes)
|
||||||
|
limiter.Reset(client)
|
||||||
|
|
||||||
|
// The client has its whole allowance of bytes again.
|
||||||
|
wantBytesCount(t, limiter, client, start, 1000, "")
|
||||||
|
}
|
||||||
|
|
||||||
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -86,14 +249,14 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
|||||||
start := midnight()
|
start := midnight()
|
||||||
|
|
||||||
for range 3 {
|
for range 3 {
|
||||||
limiter.Count(client, start)
|
limiter.Count(client, start, whole)
|
||||||
}
|
}
|
||||||
|
|
||||||
// A quarter into the next hour, the minute has only this request. The
|
// A quarter into the next hour, the minute has only this request. The
|
||||||
// hour still covers three quarters of the bucket before, with its three
|
// hour still covers three quarters of the bucket before, with its three
|
||||||
// requests, which count 2.25, and this one: 3.25. The day covers all
|
// requests, which count 2.25, and this one: 3.25. The day covers all
|
||||||
// four.
|
// four.
|
||||||
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
|
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole)
|
||||||
|
|
||||||
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
|
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
|
||||||
if counts != want {
|
if counts != want {
|
||||||
@@ -261,9 +424,24 @@ func wantCount(
|
|||||||
) {
|
) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
_, hit, _ := limiter.Count(client, now)
|
_, hit, _ := limiter.Count(client, now, whole)
|
||||||
if hit.Window != want {
|
if hit.Window != want {
|
||||||
t.Errorf("request from %s at %s is over %q, want %q",
|
t.Errorf("request from %s at %s is over %q, want %q",
|
||||||
client, now.Format(time.RFC3339), hit.Window, want)
|
client, now.Format(time.RFC3339), hit.Window, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// wantBytesCount counts bytes from client at now, and checks the kind of
|
||||||
|
// the limit they break, "" for none.
|
||||||
|
func wantBytesCount(
|
||||||
|
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, now time.Time,
|
||||||
|
bytes int64, want string,
|
||||||
|
) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
_, hit, _ := limiter.CountBytes(client, now, bytes, whole)
|
||||||
|
if hit.Kind != want {
|
||||||
|
t.Errorf("%d bytes from %s at %s break a limit on %q, want %q",
|
||||||
|
bytes, client, now.Format(time.RFC3339), hit.Kind, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ func TestSnapshotListsTheClientsByAddress(t *testing.T) {
|
|||||||
|
|
||||||
limiter := ratelimit.New(ratelimit.Limits{})
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
for _, i := range []int{2, 3, 0, 1} {
|
for _, i := range []int{2, 3, 0, 1} {
|
||||||
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
|
limiter.Count(netip.MustParsePrefix(want[i]), midnight(), whole)
|
||||||
}
|
}
|
||||||
|
|
||||||
snapshot := limiter.Snapshot()
|
snapshot := limiter.Snapshot()
|
||||||
@@ -63,7 +63,8 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
|
|||||||
start := midnight()
|
start := midnight()
|
||||||
|
|
||||||
limiter := ratelimit.New(ratelimit.Limits{})
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
limiter.Count(client, start)
|
limiter.Count(client, start, whole)
|
||||||
|
limiter.CountBytes(client, start, 5, whole)
|
||||||
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
||||||
|
|
||||||
loaded := func(now time.Time) ratelimit.Client {
|
loaded := func(now time.Time) ratelimit.Client {
|
||||||
@@ -76,19 +77,25 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Two minutes on, the window that ends then covers neither of the
|
// Two minutes on, the window that ends then covers neither of the
|
||||||
// minute's buckets, which are emptied; the hour's and the day's stay,
|
// minute's buckets, of requests and of bytes, which are emptied; the
|
||||||
// and so does the history.
|
// hour's and the day's stay, and so does the history.
|
||||||
got := loaded(start.Add(2 * time.Minute))
|
got := loaded(start.Add(2 * time.Minute))
|
||||||
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
|
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
|
||||||
got.Day.Current != 1 || got.History.Requests != 1 {
|
got.Day.Current != 1 || got.History.Requests != 1 {
|
||||||
t.Errorf("loaded two minutes on as %+v", got)
|
t.Errorf("loaded two minutes on as %+v", got)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if got.MinuteBytes != (ratelimit.Buckets{}) || got.HourBytes.Current != 5 ||
|
||||||
|
got.DayBytes.Current != 5 {
|
||||||
|
t.Errorf("loaded two minutes on with buckets of bytes %+v, %+v and %+v",
|
||||||
|
got.MinuteBytes, got.HourBytes, got.DayBytes)
|
||||||
|
}
|
||||||
|
|
||||||
// A moment before, the window still covers some of the earlier one.
|
// A moment before, the window still covers some of the earlier one.
|
||||||
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
|
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
|
||||||
if got.Minute.Current != 1 {
|
if got.Minute.Current != 1 || got.MinuteBytes.Current != 5 {
|
||||||
t.Errorf("loaded just under two minutes on with minute buckets %+v",
|
t.Errorf("loaded just under two minutes on with minute buckets %+v and %+v",
|
||||||
got.Minute)
|
got.Minute, got.MinuteBytes)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// OffenceLimit is the offence a request line names for a request that
|
// OffenceLimit is the offence a request line names for a request that
|
||||||
// broke a rate limit.
|
// broke a rate limit, or whose bytes broke a byte limit.
|
||||||
const OffenceLimit = "limit"
|
const OffenceLimit = "limit"
|
||||||
|
|
||||||
// timeLayout is RFC 3339 with milliseconds.
|
// timeLayout is RFC 3339 with milliseconds.
|
||||||
@@ -80,11 +80,14 @@ type Line struct {
|
|||||||
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
|
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
|
||||||
// or a new one, and is sent on to the app. ForwardedFor is the
|
// or a new one, and is sent on to the app. ForwardedFor is the
|
||||||
// X-Forwarded-For header as received. ClientGroup is the netblock the
|
// X-Forwarded-For header as received. ClientGroup is the netblock the
|
||||||
// client is counted as.
|
// client is counted as. ASN, ASName and Country are the client's AS
|
||||||
|
// number, AS name and country, as looked up.
|
||||||
RequestID string `json:"request_id"`
|
RequestID string `json:"request_id"`
|
||||||
PeerIP string `json:"peer_ip"`
|
PeerIP string `json:"peer_ip"`
|
||||||
ForwardedFor string `json:"forwarded_for,omitempty"`
|
ForwardedFor string `json:"forwarded_for,omitempty"`
|
||||||
ClientGroup string `json:"client_group"`
|
ClientGroup string `json:"client_group"`
|
||||||
|
ASN string `json:"asn"`
|
||||||
|
ASName string `json:"as_name"`
|
||||||
Country string `json:"country"`
|
Country string `json:"country"`
|
||||||
ContentType string `json:"content_type,omitempty"`
|
ContentType string `json:"content_type,omitempty"`
|
||||||
// ContentLength is the length of its body the request announced.
|
// ContentLength is the length of its body the request announced.
|
||||||
@@ -114,13 +117,25 @@ type Line struct {
|
|||||||
// ActionBanned, ActionCountryDenied, ActionRateLimited or
|
// ActionBanned, ActionCountryDenied, ActionRateLimited or
|
||||||
// ActionRuleBlocked.
|
// ActionRuleBlocked.
|
||||||
WouldAction string `json:"would_action,omitempty"`
|
WouldAction string `json:"would_action,omitempty"`
|
||||||
// Counts are the client's requests as the rate limits counted them
|
// LimitPercent and LimitPercentSetting are, for a request the rate
|
||||||
// with this one, for a request they counted.
|
// limits counted whose client a biased threshold gives a percentage of
|
||||||
|
// the rate limits below 100, that percentage and the setting that gave
|
||||||
|
// it. BytesPercent and BytesPercentSetting are the same for the byte
|
||||||
|
// limits.
|
||||||
|
LimitPercent *int64 `json:"limit_percent,omitempty"`
|
||||||
|
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
|
||||||
|
BytesPercent *int64 `json:"bytes_percent,omitempty"`
|
||||||
|
BytesPercentSetting string `json:"bytes_percent_setting,omitempty"`
|
||||||
|
// Counts are, for a request the rate limits counted, the client's
|
||||||
|
// requests as they counted them with this one, and its bytes as the
|
||||||
|
// byte limits counted them, with this request's once it has ended if
|
||||||
|
// they count them.
|
||||||
Counts ratelimit.Counts `json:"counts,omitzero"`
|
Counts ratelimit.Counts `json:"counts,omitzero"`
|
||||||
// RuleIDs are the ids of the rule file rules the request matched.
|
// RuleIDs are the ids of the rule file rules the request matched.
|
||||||
RuleIDs []string `json:"rule_ids,omitempty"`
|
RuleIDs []string `json:"rule_ids,omitempty"`
|
||||||
// LimitHit is the window whose rate limit the request went over:
|
// LimitHit is the window whose limit the request went over, named as
|
||||||
// minute, hour or day.
|
// Counts names its count: minute, hour or day for a rate limit, and
|
||||||
|
// minute_bytes, hour_bytes or day_bytes for a byte limit.
|
||||||
LimitHit string `json:"limit_hit,omitempty"`
|
LimitHit string `json:"limit_hit,omitempty"`
|
||||||
// Offence is the offence the request was held as, OffenceLimit.
|
// Offence is the offence the request was held as, OffenceLimit.
|
||||||
Offence string `json:"offence,omitempty"`
|
Offence string `json:"offence,omitempty"`
|
||||||
@@ -171,8 +186,8 @@ func Milliseconds(d time.Duration) float64 {
|
|||||||
|
|
||||||
// NewProcessLogger returns the logger for the process's own messages:
|
// NewProcessLogger returns the logger for the process's own messages:
|
||||||
// JSON lines on w, marked "type":"process", with the time in the same form
|
// JSON lines on w, marked "type":"process", with the time in the same form
|
||||||
// as a request line's.
|
// as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance.
|
||||||
func NewProcessLogger(w io.Writer) *slog.Logger {
|
func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
|
||||||
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
||||||
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
||||||
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
||||||
@@ -183,5 +198,5 @@ func NewProcessLogger(w io.Writer) *slog.Logger {
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
return slog.New(handler).With("type", "process")
|
return slog.New(handler).With("type", "process", "instance", instanceName)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -65,12 +65,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestProcessLinesAreMarkedProcess(t *testing.T) {
|
func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
var out bytes.Buffer
|
var out bytes.Buffer
|
||||||
|
|
||||||
requestlog.NewProcessLogger(&out).Info("starting", "version", "v1")
|
requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1")
|
||||||
|
|
||||||
var fields map[string]any
|
var fields map[string]any
|
||||||
|
|
||||||
@@ -79,8 +79,9 @@ func TestProcessLinesAreMarkedProcess(t *testing.T) {
|
|||||||
t.Fatalf("decode %q: %v", out.String(), err)
|
t.Fatalf("decode %q: %v", out.String(), err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if fields["type"] != "process" || fields["msg"] != "starting" ||
|
if fields["type"] != "process" || fields["instance"] != "fsn1app1/gitea" ||
|
||||||
fields["level"] != "INFO" || fields["version"] != "v1" {
|
fields["msg"] != "starting" || fields["level"] != "INFO" ||
|
||||||
|
fields["version"] != "v1" {
|
||||||
t.Errorf("process line %v", fields)
|
t.Errorf("process line %v", fields)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+28
-14
@@ -22,6 +22,7 @@ import (
|
|||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
"github.com/fsnotify/fsnotify"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -100,6 +101,8 @@ type Params struct {
|
|||||||
// ProcessLog receives how many rules were read, and the error in a
|
// ProcessLog receives how many rules were read, and the error in a
|
||||||
// rule file edited while smallwebwaf runs.
|
// rule file edited while smallwebwaf runs.
|
||||||
ProcessLog *slog.Logger
|
ProcessLog *slog.Logger
|
||||||
|
// Alerts receive a file_error alert for that error.
|
||||||
|
Alerts *alerts.Queue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Files are the rule files of a running smallwebwaf, and the rules read
|
// Files are the rule files of a running smallwebwaf, and the rules read
|
||||||
@@ -126,7 +129,7 @@ func Load(params Params) (*Files, error) {
|
|||||||
return f, nil
|
return f, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
rules, err := read(params.Dir)
|
rules, _, err := read(params.Dir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -223,13 +226,21 @@ func (f *Files) readAfterChanges(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// readAgain reads the rule files again, in place of the rules loaded, or
|
// readAgain reads the rule files again, in place of the rules loaded, or
|
||||||
// logs the error that keeps the rules as they were.
|
// logs the error that keeps the rules as they were, and raises a
|
||||||
|
// file_error alert for it, for the file it is in.
|
||||||
func (f *Files) readAgain() {
|
func (f *Files) readAgain() {
|
||||||
rules, err := read(f.params.Dir)
|
rules, path, err := read(f.params.Dir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.params.ProcessLog.Error(
|
const kept = "a rule file has an error, and the rules stay as they were"
|
||||||
"a rule file has an error, and the rules stay as they were",
|
|
||||||
"error", err.Error())
|
// Raised before it is logged, so that the alert is there once the
|
||||||
|
// log line is.
|
||||||
|
f.params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventFileError,
|
||||||
|
Reason: kept,
|
||||||
|
Detail: map[string]any{"file": path, "error": err.Error()},
|
||||||
|
})
|
||||||
|
f.params.ProcessLog.Error(kept, "error", err.Error())
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -246,13 +257,14 @@ func (f *Files) logRead(count int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// read returns the rules of every rule file in dir, in the order of the
|
// read returns the rules of every rule file in dir, in the order of the
|
||||||
// files' names, and then of their lines. A file whose name starts with a
|
// files' names, and then of their lines, or an error, with the path of the
|
||||||
// dot, such as an editor's lock file .#50-app.rules, is not a rule file,
|
// rule file it is in, or dir. A file whose name starts with a dot, such as
|
||||||
// as a shell's *.rules would not match it.
|
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's
|
||||||
func read(dir string) ([]Rule, error) {
|
// *.rules would not match it.
|
||||||
|
func read(dir string) ([]Rule, string, error) {
|
||||||
entries, err := os.ReadDir(dir)
|
entries, err := os.ReadDir(dir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
|
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var rules []Rule
|
var rules []Rule
|
||||||
@@ -266,13 +278,15 @@ func read(dir string) ([]Rule, error) {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
rules, err = readFile(filepath.Join(dir, name), rules, places)
|
path := filepath.Join(dir, name)
|
||||||
|
|
||||||
|
rules, err = readFile(path, rules, places)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, path, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return rules, nil
|
return rules, "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// readFile appends the rules of the rule file at path to rules. places
|
// readFile appends the rules of the rule file at path to rules. places
|
||||||
|
|||||||
@@ -7,12 +7,15 @@ import (
|
|||||||
"maps"
|
"maps"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"slices"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -337,7 +340,7 @@ func TestEditsTakenInWhileRunning(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
||||||
files, lines := watch(t, dir)
|
files, lines, _ := watch(t, dir)
|
||||||
|
|
||||||
// matches reports whether path matches a rule.
|
// matches reports whether path matches a rule.
|
||||||
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
|
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
|
||||||
@@ -366,7 +369,7 @@ func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
||||||
files, lines := watch(t, dir)
|
files, lines, queue := watch(t, dir)
|
||||||
|
|
||||||
// The edit's second line has an unknown action, so the rules stay as
|
// The edit's second line has an unknown action, so the rules stay as
|
||||||
// they were, the first line's earlier version included.
|
// they were, the first line's earlier version included.
|
||||||
@@ -380,13 +383,27 @@ func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
|
|||||||
t.Errorf("logged %v, want an error %q", line, want)
|
t.Errorf("logged %v, want an error %q", line, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The error is raised as a file_error alert too, for the file.
|
||||||
|
wantFileError := func() {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
|
||||||
|
waiting[0].Reason != hasError || waiting[0].Detail["error"] != want ||
|
||||||
|
waiting[0].Detail["file"] != filepath.Join(dir, firstFile) {
|
||||||
|
t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wantFileError()
|
||||||
|
|
||||||
wantMatched(t, files, get(t, "/first"), "first")
|
wantMatched(t, files, get(t, "/first"), "first")
|
||||||
wantMatched(t, files, get(t, "/second"))
|
wantMatched(t, files, get(t, "/second"))
|
||||||
|
|
||||||
// Once mended, the file is read again.
|
// Once mended, the file is read again, and raises no alert.
|
||||||
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
|
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
|
||||||
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
|
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
|
||||||
wantMatched(t, files, get(t, "/edited"), "first")
|
wantMatched(t, files, get(t, "/edited"), "first")
|
||||||
|
wantFileError()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
|
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
|
||||||
@@ -504,7 +521,8 @@ func save(t *testing.T, dir, name, content string) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// newParams returns Params for the rule files in dir, switched on, with
|
// newParams returns Params for the rule files in dir, switched on, with
|
||||||
// the process log in the processLog returned.
|
// the process log in the processLog returned, and the alerts waiting in a
|
||||||
|
// queue for a webhook that is never sent them.
|
||||||
func newParams(dir string) (rules.Params, processLog) {
|
func newParams(dir string) (rules.Params, processLog) {
|
||||||
lines := make(processLog, maxLogLines)
|
lines := make(processLog, maxLogLines)
|
||||||
|
|
||||||
@@ -512,6 +530,12 @@ func newParams(dir string) (rules.Params, processLog) {
|
|||||||
Dir: dir,
|
Dir: dir,
|
||||||
Enabled: true,
|
Enabled: true,
|
||||||
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
|
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
|
||||||
|
Alerts: alerts.New(alerts.Params{
|
||||||
|
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||||
|
Events: alerts.Events(),
|
||||||
|
Cooldown: 15 * time.Minute,
|
||||||
|
Now: time.Now,
|
||||||
|
}),
|
||||||
}, lines
|
}, lines
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -531,8 +555,9 @@ func load(t *testing.T, files ruleFiles) *rules.Files {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// watch loads the rules in dir, runs their Watch until the test ends, and
|
// watch loads the rules in dir, runs their Watch until the test ends, and
|
||||||
// waits until it watches the directory.
|
// waits until it watches the directory. It returns the alerts' queue as
|
||||||
func watch(t *testing.T, dir string) (*rules.Files, processLog) {
|
// well.
|
||||||
|
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
params, lines := newParams(dir)
|
params, lines := newParams(dir)
|
||||||
@@ -557,7 +582,7 @@ func watch(t *testing.T, dir string) (*rules.Files, processLog) {
|
|||||||
|
|
||||||
lines.waitFor(t, watching)
|
lines.waitFor(t, watching)
|
||||||
|
|
||||||
return files, lines
|
return files, lines, params.Alerts
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantRefused checks that loading the rule files in dir fails with the
|
// wantRefused checks that loading the rule files in dir fails with the
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
"github.com/fsnotify/fsnotify"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
)
|
)
|
||||||
|
|
||||||
// The tests below run readAfterChanges in a synctest bubble, where time is
|
// The tests below run readAfterChanges in a synctest bubble, where time is
|
||||||
@@ -91,6 +93,7 @@ func load(t *testing.T, dir string) *Files {
|
|||||||
|
|
||||||
files, err := Load(Params{
|
files, err := Load(Params{
|
||||||
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
|
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Alerts: alerts.New(alerts.Params{}),
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("load: %v", err)
|
t.Fatalf("load: %v", err)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
||||||
// the rule files and the state files, serves requests until it is told to
|
// the rule files, the lookup database and the state files, serves requests
|
||||||
// stop, and then stops in an orderly way, writing the state files.
|
// until it is told to stop, and then stops in an orderly way, writing the
|
||||||
|
// state files.
|
||||||
package smallwebwaf
|
package smallwebwaf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -15,6 +16,7 @@ import (
|
|||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
@@ -63,11 +65,12 @@ func Main(version string) int {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run reads the settings, the rule files and the state files, then serves
|
// Run reads the settings, the rule files, the lookup database and the
|
||||||
// requests until ctx is done. It returns the process's exit status, 1
|
// state files, then serves requests until ctx is done. It returns the
|
||||||
// when smallwebwaf cannot start.
|
// process's exit status, 1 when smallwebwaf cannot start.
|
||||||
func Run(ctx context.Context, params Params) int {
|
func Run(ctx context.Context, params Params) int {
|
||||||
processLog := requestlog.NewProcessLogger(params.Stdout)
|
processLog := requestlog.NewProcessLogger(params.Stdout,
|
||||||
|
config.InstanceName(params.LookupEnv))
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(params.LookupEnv)
|
cfg, err := config.FromEnvironment(params.LookupEnv)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -85,16 +88,22 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
if cfg.LogRemoteURL != nil {
|
if cfg.LogRemoteURL != nil {
|
||||||
remote = newRemoteLogSender(cfg)
|
remote = newRemoteLogSender(cfg)
|
||||||
stdout = io.MultiWriter(params.Stdout, remote)
|
stdout = io.MultiWriter(params.Stdout, remote)
|
||||||
processLog = requestlog.NewProcessLogger(stdout)
|
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName)
|
||||||
|
|
||||||
stopSending := startSending(ctx, remote, processLog)
|
stopSending := startSending(ctx, remote, processLog)
|
||||||
defer stopSending()
|
defer stopSending()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The state files and the alerts give times in UTC.
|
||||||
|
now := func() time.Time { return time.Now().UTC() }
|
||||||
|
|
||||||
|
alertQueue := newAlertQueue(cfg, now, processLog)
|
||||||
|
|
||||||
ruleFiles, err := rules.Load(rules.Params{
|
ruleFiles, err := rules.Load(rules.Params{
|
||||||
Dir: cfg.RulesDir,
|
Dir: cfg.RulesDir,
|
||||||
Enabled: cfg.RulesEnabled,
|
Enabled: cfg.RulesEnabled,
|
||||||
ProcessLog: processLog,
|
ProcessLog: processLog,
|
||||||
|
Alerts: alertQueue,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
processLog.Error("cannot use the rule files", "error", err.Error())
|
processLog.Error("cannot use the rule files", "error", err.Error())
|
||||||
@@ -102,32 +111,18 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
// The state files give times in UTC.
|
server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue)
|
||||||
now := func() time.Time { return time.Now().UTC() }
|
if err != nil {
|
||||||
|
processLog.Error("cannot use the lookup database", "error", err.Error())
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: stdout,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: lookup.URL,
|
|
||||||
Now: now,
|
|
||||||
Rules: ruleFiles,
|
|
||||||
})
|
|
||||||
if remote != nil {
|
if remote != nil {
|
||||||
server.Metrics.AddRemoteLog(remote)
|
server.Metrics.AddRemoteLog(remote)
|
||||||
}
|
}
|
||||||
|
|
||||||
files, err := state.Load(state.Params{
|
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
|
||||||
Dir: cfg.StateDir,
|
|
||||||
WriteDelay: cfg.StateWriteDelay,
|
|
||||||
CounterInterval: cfg.StateCounterInterval,
|
|
||||||
Ledger: server.Ledger,
|
|
||||||
Limiter: server.Limiter,
|
|
||||||
GeoJS: server.GeoJS,
|
|
||||||
Now: now,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
Metrics: server.Metrics,
|
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
processLog.Error("cannot use the state files", "error", err.Error())
|
processLog.Error("cannot use the state files", "error", err.Error())
|
||||||
|
|
||||||
@@ -147,7 +142,89 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
"address", listener.Addr().String(),
|
"address", listener.Addr().String(),
|
||||||
"settings", cfg)
|
"settings", cfg)
|
||||||
|
|
||||||
return serve(ctx, server.Server, listener, files, ruleFiles, processLog)
|
return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog)
|
||||||
|
}
|
||||||
|
|
||||||
|
// newServer returns the server smallwebwaf runs, with the metrics of the
|
||||||
|
// alerts, after reading the lookup database while SWWAF_LOOKUP_SOURCE is
|
||||||
|
// file. A lookup database that cannot be read is an error.
|
||||||
|
func newServer(
|
||||||
|
cfg *config.Config, stdout io.Writer, processLog *slog.Logger,
|
||||||
|
now func() time.Time, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
||||||
|
) (*proxy.Server, error) {
|
||||||
|
var lookupFile *lookup.File
|
||||||
|
|
||||||
|
if cfg.LookupSource == "file" {
|
||||||
|
var err error
|
||||||
|
|
||||||
|
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
||||||
|
Path: cfg.LookupDBPath,
|
||||||
|
Now: now,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
Alerts: alertQueue,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
server := proxy.New(proxy.Params{
|
||||||
|
Config: cfg,
|
||||||
|
RequestLog: stdout,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
GeoJSURL: lookup.URL,
|
||||||
|
LookupFile: lookupFile,
|
||||||
|
Now: now,
|
||||||
|
Rules: ruleFiles,
|
||||||
|
Alerts: alertQueue,
|
||||||
|
})
|
||||||
|
server.Metrics.AddAlerts(alertQueue)
|
||||||
|
|
||||||
|
if lookupFile != nil {
|
||||||
|
server.Metrics.AddLookupFile(lookupFile.LastRead, lookupFile.ReadFailures)
|
||||||
|
}
|
||||||
|
|
||||||
|
return server, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// loadStateFiles reads the state files into the parts of server and into
|
||||||
|
// alertQueue, as state.Load does.
|
||||||
|
func loadStateFiles(
|
||||||
|
cfg *config.Config, server *proxy.Server, alertQueue *alerts.Queue,
|
||||||
|
now func() time.Time, processLog *slog.Logger,
|
||||||
|
) (*state.Files, error) {
|
||||||
|
return state.Load(state.Params{
|
||||||
|
Dir: cfg.StateDir,
|
||||||
|
WriteDelay: cfg.StateWriteDelay,
|
||||||
|
CounterInterval: cfg.StateCounterInterval,
|
||||||
|
Ledger: server.Ledger,
|
||||||
|
Limiter: server.Limiter,
|
||||||
|
GeoJS: server.GeoJS,
|
||||||
|
Alerts: alertQueue,
|
||||||
|
Now: now,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
Metrics: server.Metrics,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// newAlertQueue returns the queue of the alerts to the webhook, Slack and
|
||||||
|
// ntfy, with the settings for them.
|
||||||
|
func newAlertQueue(
|
||||||
|
cfg *config.Config, now func() time.Time, processLog *slog.Logger,
|
||||||
|
) *alerts.Queue {
|
||||||
|
return alerts.New(alerts.Params{
|
||||||
|
WebhookURL: cfg.AlertWebhookURL,
|
||||||
|
WebhookHeaders: cfg.AlertWebhookHeaders,
|
||||||
|
SlackURL: cfg.AlertSlackWebhookURL,
|
||||||
|
NtfyURL: cfg.AlertNtfyURL,
|
||||||
|
NtfyToken: cfg.AlertNtfyToken,
|
||||||
|
Events: cfg.AlertEvents,
|
||||||
|
Cooldown: cfg.AlertCooldown,
|
||||||
|
MaxPerHour: cfg.AlertMaxPerHour,
|
||||||
|
Instance: cfg.InstanceName,
|
||||||
|
Now: now,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// newRemoteLogSender returns a sender of the log lines to
|
// newRemoteLogSender returns a sender of the log lines to
|
||||||
@@ -188,12 +265,15 @@ func startSending(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// serve serves requests on listener, writes the state files as they are
|
// serve serves requests on listener, writes the state files as they are
|
||||||
// due, takes in an admin's edits of them, and reads the rule files again
|
// due, takes in an admin's edits of them, reads the rule files again as
|
||||||
// as they change, until ctx is done. Then it gives the requests in
|
// they change, and the lookup database when it is replaced, and sends the
|
||||||
// progress shutdownTimeout to finish, and writes every state file.
|
// alerts, until ctx is done. Then it gives the requests in progress
|
||||||
|
// shutdownTimeout to finish, and writes every state file, alerts.json with
|
||||||
|
// the alerts still waiting.
|
||||||
func serve(
|
func serve(
|
||||||
ctx context.Context, server *http.Server, listener net.Listener,
|
ctx context.Context, server *proxy.Server, listener net.Listener,
|
||||||
files *state.Files, ruleFiles *rules.Files, processLog *slog.Logger,
|
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
||||||
|
processLog *slog.Logger,
|
||||||
) int {
|
) int {
|
||||||
served := make(chan error, 1)
|
served := make(chan error, 1)
|
||||||
|
|
||||||
@@ -204,24 +284,15 @@ func serve(
|
|||||||
writing, stopWriting := context.WithCancel(ctx)
|
writing, stopWriting := context.WithCancel(ctx)
|
||||||
defer stopWriting()
|
defer stopWriting()
|
||||||
|
|
||||||
written := make(chan struct{})
|
written := inBackground(func() { files.Run(writing) })
|
||||||
watched := make(chan struct{})
|
watched := inBackground(func() { files.Watch(writing) })
|
||||||
rulesWatched := make(chan struct{})
|
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
|
||||||
|
lookupFileWatched := inBackground(func() {
|
||||||
go func() {
|
if server.LookupFile != nil {
|
||||||
files.Run(writing)
|
server.LookupFile.Watch(writing)
|
||||||
close(written)
|
}
|
||||||
}()
|
})
|
||||||
|
alertsSent := inBackground(func() { alertQueue.Run(writing) })
|
||||||
go func() {
|
|
||||||
files.Watch(writing)
|
|
||||||
close(watched)
|
|
||||||
}()
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
ruleFiles.Watch(writing)
|
|
||||||
close(rulesWatched)
|
|
||||||
}()
|
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-served:
|
case err := <-served:
|
||||||
@@ -253,7 +324,8 @@ func serve(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Run and Watch have ended, so nothing else reads or writes the
|
// Run and Watch have ended, so nothing else reads or writes the
|
||||||
// files. Every request has ended too, but for two kinds
|
// files, and no alert is being sent, so that alerts.json keeps every
|
||||||
|
// alert not yet sent. Every request has ended too, but for two kinds
|
||||||
// that Go's server does not wait for: one cut off because Shutdown
|
// that Go's server does not wait for: one cut off because Shutdown
|
||||||
// timed out, and one whose connection switched protocols, such as a
|
// timed out, and one whose connection switched protocols, such as a
|
||||||
// WebSocket. Such a request adds to its client's history only as it
|
// WebSocket. Such a request adds to its client's history only as it
|
||||||
@@ -262,6 +334,8 @@ func serve(
|
|||||||
<-written
|
<-written
|
||||||
<-watched
|
<-watched
|
||||||
<-rulesWatched
|
<-rulesWatched
|
||||||
|
<-lookupFileWatched
|
||||||
|
<-alertsSent
|
||||||
|
|
||||||
err = files.WriteAll()
|
err = files.WriteAll()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -274,3 +348,16 @@ func serve(
|
|||||||
|
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// inBackground runs task on a goroutine of its own, and returns a channel
|
||||||
|
// that is closed once task has returned.
|
||||||
|
func inBackground(task func()) <-chan struct{} {
|
||||||
|
done := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
task()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return done
|
||||||
|
}
|
||||||
|
|||||||
@@ -14,9 +14,11 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
|
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -37,11 +39,20 @@ const (
|
|||||||
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
||||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
rulesDir = "SWWAF_RULES_DIR"
|
rulesDir = "SWWAF_RULES_DIR"
|
||||||
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
|
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
||||||
|
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
||||||
|
instanceName = "SWWAF_INSTANCE_NAME"
|
||||||
|
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
|
||||||
|
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
|
||||||
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set.
|
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set.
|
||||||
adminSecret = "fedcba9876543210fedcba9876543210"
|
adminSecret = "fedcba9876543210fedcba9876543210"
|
||||||
|
// instance is the SWWAF_INSTANCE_NAME the tests set where they look at
|
||||||
|
// it.
|
||||||
|
instance = "fsn1app1/gitea"
|
||||||
// greeting is what the tests' app answers.
|
// greeting is what the tests' app answers.
|
||||||
greeting = "hello from the app"
|
greeting = "hello from the app"
|
||||||
|
// placed is the client the tests' lookup databases place.
|
||||||
|
placed = "203.0.113.9"
|
||||||
)
|
)
|
||||||
|
|
||||||
// output collects what smallwebwaf writes on stdout.
|
// output collects what smallwebwaf writes on stdout.
|
||||||
@@ -98,12 +109,16 @@ func (o *output) text() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// run runs smallwebwaf with the settings in env until ctx is done, and
|
// run runs smallwebwaf with the settings in env until ctx is done, and
|
||||||
// returns its exit status.
|
// returns its exit status. SWWAF_LOOKUP_SOURCE is off unless env sets it,
|
||||||
|
// so that no test sends GeoJS its clients' addresses.
|
||||||
func run(ctx context.Context, env map[string]string, out *output) int {
|
func run(ctx context.Context, env map[string]string, out *output) int {
|
||||||
return smallwebwaf.Run(ctx, smallwebwaf.Params{
|
return smallwebwaf.Run(ctx, smallwebwaf.Params{
|
||||||
Version: testVersion,
|
Version: testVersion,
|
||||||
LookupEnv: func(name string) (string, bool) {
|
LookupEnv: func(name string) (string, bool) {
|
||||||
value, ok := env[name]
|
value, ok := env[name]
|
||||||
|
if !ok && name == "SWWAF_LOOKUP_SOURCE" {
|
||||||
|
return "off", true
|
||||||
|
}
|
||||||
|
|
||||||
return value, ok
|
return value, ok
|
||||||
},
|
},
|
||||||
@@ -116,7 +131,10 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
|
|||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
|
|
||||||
status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out)
|
status := run(t.Context(), map[string]string{
|
||||||
|
"SWWAF_REQUEST_MAX_BYTES": "lots",
|
||||||
|
instanceName: instance,
|
||||||
|
}, out)
|
||||||
if status != 1 {
|
if status != 1 {
|
||||||
t.Errorf("exit status %d, want 1", status)
|
t.Errorf("exit status %d, want 1", status)
|
||||||
}
|
}
|
||||||
@@ -124,7 +142,8 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
|
|||||||
line := out.line(t, "msg", "invalid setting")
|
line := out.line(t, "msg", "invalid setting")
|
||||||
message, _ := line["error"].(string)
|
message, _ := line["error"].(string)
|
||||||
|
|
||||||
if line["type"] != "process" || line["level"] != "ERROR" ||
|
if line["type"] != "process" || line["instance"] != instance ||
|
||||||
|
line["level"] != "ERROR" ||
|
||||||
!strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") {
|
!strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") {
|
||||||
t.Errorf("start refused with %v", line)
|
t.Errorf("start refused with %v", line)
|
||||||
}
|
}
|
||||||
@@ -135,7 +154,7 @@ func TestShortTokenStopsTheStartUnshown(t *testing.T) {
|
|||||||
|
|
||||||
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
|
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
|
||||||
|
|
||||||
for _, name := range []string{adminToken, "SWWAF_METRICS_TOKEN"} {
|
for _, name := range []string{adminToken, metricsToken} {
|
||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -229,6 +248,46 @@ func TestServesUntilToldToStop(t *testing.T) {
|
|||||||
out.line(t, "msg", "stopped")
|
out.line(t, "msg", "stopped")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestEveryLogLineAndMetricCarriesTheInstanceName(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const token = "0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
metricsToken: token,
|
||||||
|
instanceName: instance,
|
||||||
|
}
|
||||||
|
|
||||||
|
var metrics string
|
||||||
|
|
||||||
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
|
metrics = metricsText(t, url+"_smallwebwaf/metrics", token)
|
||||||
|
})
|
||||||
|
|
||||||
|
// The process's lines from its start to its stop, and the request's.
|
||||||
|
for line := range strings.Lines(out.text()) {
|
||||||
|
var fields map[string]any
|
||||||
|
|
||||||
|
err := json.Unmarshal([]byte(line), &fields)
|
||||||
|
if err != nil || fields["instance"] != instance {
|
||||||
|
t.Errorf("line %q (%v), want instance %s", line, err, instance)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Each series, Go's and the process's included; the other lines are
|
||||||
|
// the comments.
|
||||||
|
for line := range strings.Lines(metrics) {
|
||||||
|
if !strings.HasPrefix(line, "#") &&
|
||||||
|
!strings.Contains(line, `instance="fsn1app1/gitea"`) {
|
||||||
|
t.Errorf("series %q, without instance=\"fsn1app1/gitea\"", line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestStateKeptAcrossRestarts(t *testing.T) {
|
func TestStateKeptAcrossRestarts(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -441,6 +500,107 @@ func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
|
|||||||
": no such file or directory")
|
": no such file or directory")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLookupDatabaseReplacedWhileRunningTakesEffect(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const token = "0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
writeLookupDatabase(t, path, "DE")
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
trustedProxies: localhost + "/32",
|
||||||
|
metricsToken: token,
|
||||||
|
instanceName: instance,
|
||||||
|
lookupSource: "file",
|
||||||
|
lookupDBPath: path,
|
||||||
|
"SWWAF_DENIED_COUNTRIES": "kp",
|
||||||
|
// The requests sent until a replacement takes effect, and those
|
||||||
|
// for the metrics, must not break a rate limit, whose ban would
|
||||||
|
// refuse them too.
|
||||||
|
"SWWAF_RATE_LIMIT_EXEMPT_NETS": placed + "," + localhost,
|
||||||
|
}
|
||||||
|
began := time.Now()
|
||||||
|
// Each replacement is written beside the file and renamed over it, as
|
||||||
|
// a refresh is.
|
||||||
|
replacement := path + ".new"
|
||||||
|
|
||||||
|
var lastRead float64
|
||||||
|
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
wantStatus(t, url, placed, http.StatusOK)
|
||||||
|
|
||||||
|
// smallwebwaf reads the file again once it watches its directory,
|
||||||
|
// which may be after the first replacement. Once that has taken
|
||||||
|
// effect, only the watch can show the next. Each takes as long as
|
||||||
|
// it takes, so that a slow test process cannot fail the test.
|
||||||
|
writeLookupDatabase(t, replacement, "KP")
|
||||||
|
rename(t, replacement, path)
|
||||||
|
|
||||||
|
for statusFrom(t, url, placed) != http.StatusForbidden {
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
|
||||||
|
writeLookupDatabase(t, replacement, "DE")
|
||||||
|
rename(t, replacement, path)
|
||||||
|
|
||||||
|
for statusFrom(t, url, placed) != http.StatusOK {
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
|
||||||
|
// One that cannot be read leaves the file read before in use.
|
||||||
|
err := os.WriteFile(replacement, []byte("not a lookup database\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", replacement, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rename(t, replacement, path)
|
||||||
|
|
||||||
|
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token,
|
||||||
|
`smallwebwaf_lookup_database_read_failures_total{instance="fsn1app1/gitea"} 1`)
|
||||||
|
lastRead = seriesValue(t, metrics,
|
||||||
|
`smallwebwaf_lookup_database_last_read_timestamp_seconds{instance="fsn1app1/gitea"}`)
|
||||||
|
|
||||||
|
wantStatus(t, url, placed, http.StatusOK)
|
||||||
|
})
|
||||||
|
|
||||||
|
if lastRead < float64(began.Unix()) || lastRead > float64(time.Now().Unix()) {
|
||||||
|
t.Errorf("the lookup database was last read at %v, want a time in the test", lastRead)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLookupDatabaseThatCannotBeReadStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||||
|
|
||||||
|
// If it starts instead, it is stopped after waitLimit.
|
||||||
|
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
out := &output{}
|
||||||
|
|
||||||
|
status := run(ctx, map[string]string{
|
||||||
|
listenAddr: localhost + ":0", stateDir: t.TempDir(), rulesDir: t.TempDir(),
|
||||||
|
lookupSource: "file", lookupDBPath: path,
|
||||||
|
}, out)
|
||||||
|
if status != 1 {
|
||||||
|
t.Fatalf("exit status %d, want 1", status)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := "SWWAF_LOOKUP_DB_PATH cannot be read: open " + path +
|
||||||
|
": no such file or directory"
|
||||||
|
|
||||||
|
line := out.line(t, "msg", "cannot use the lookup database")
|
||||||
|
if line["error"] != want {
|
||||||
|
t.Errorf("start refused with %v, want the error %q", line, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
|
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -516,7 +676,8 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
|
|||||||
rulesDir: t.TempDir(),
|
rulesDir: t.TempDir(),
|
||||||
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
|
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
|
||||||
"SWWAF_LOG_REMOTE_BUFFER": "1",
|
"SWWAF_LOG_REMOTE_BUFFER": "1",
|
||||||
"SWWAF_METRICS_TOKEN": token,
|
metricsToken: token,
|
||||||
|
instanceName: instance,
|
||||||
}
|
}
|
||||||
|
|
||||||
out := runUntilStopped(t, env, func(url string) {
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
@@ -526,16 +687,17 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
|
|||||||
// last.
|
// last.
|
||||||
metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
|
metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
|
||||||
for _, series := range []string{
|
for _, series := range []string{
|
||||||
"smallwebwaf_remote_log_lines_sent_total 0",
|
`smallwebwaf_remote_log_lines_sent_total{instance="fsn1app1/gitea"} 0`,
|
||||||
"smallwebwaf_remote_log_buffer_depth 1",
|
`smallwebwaf_remote_log_buffer_depth{instance="fsn1app1/gitea"} 1`,
|
||||||
} {
|
} {
|
||||||
if !strings.Contains(metrics, "\n"+series+"\n") {
|
if !strings.Contains(metrics, "\n"+series+"\n") {
|
||||||
t.Errorf("no %q in the metrics:\n%s", series, metrics)
|
t.Errorf("no %q in the metrics:\n%s", series, metrics)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") ||
|
const dropped = "\nsmallwebwaf_remote_log_lines_dropped_total" +
|
||||||
!strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") {
|
`{instance="fsn1app1/gitea"} `
|
||||||
|
if strings.Contains(metrics, dropped+"0\n") || !strings.Contains(metrics, dropped) {
|
||||||
t.Errorf("no line dropped in the metrics:\n%s", metrics)
|
t.Errorf("no line dropped in the metrics:\n%s", metrics)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -545,6 +707,170 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
out.line(t, "type", "request")
|
out.line(t, "type", "request")
|
||||||
|
|
||||||
|
// While the lines are sent, the process's lines give the instance name
|
||||||
|
// too.
|
||||||
|
line := out.line(t, "msg", "starting")
|
||||||
|
if line["instance"] != instance {
|
||||||
|
t.Errorf("start logged with instance %v, want %s", line["instance"], instance)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
webhook := startWebhook(t)
|
||||||
|
rules := t.TempDir()
|
||||||
|
|
||||||
|
err := os.WriteFile(filepath.Join(rules, "50-app.rules"),
|
||||||
|
[]byte(`probe path ban ^/\.env$`+"\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write the rule file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: dir,
|
||||||
|
rulesDir: rules,
|
||||||
|
"SWWAF_ALERT_WEBHOOK_URL": webhook.url,
|
||||||
|
"SWWAF_ALERT_WEBHOOK_HEADERS": "Authorization:Bearer " + adminSecret,
|
||||||
|
instanceName: instance,
|
||||||
|
}
|
||||||
|
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
// The probe bans the client, and the webhook is sent the alert.
|
||||||
|
wantRefused(t, url+".env")
|
||||||
|
|
||||||
|
post := webhook.waitFor(t, "ban", true)
|
||||||
|
if post.alert["client"] != localhost || post.alert["netblock"] != localhost+"/32" ||
|
||||||
|
post.authorization != "Bearer "+adminSecret {
|
||||||
|
t.Errorf("the webhook was sent %v, with Authorization %q", post.alert,
|
||||||
|
post.authorization)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The webhook fails, so the alert for the ban made permanent by the
|
||||||
|
// client's next request waits.
|
||||||
|
webhook.failing.Store(true)
|
||||||
|
wantRefused(t, url)
|
||||||
|
webhook.waitFor(t, "permanent_ban", false)
|
||||||
|
})
|
||||||
|
|
||||||
|
// alerts.json keeps it as smallwebwaf stops, and once started again,
|
||||||
|
// smallwebwaf sends it.
|
||||||
|
var file struct {
|
||||||
|
Waiting map[string][]struct {
|
||||||
|
Event string `json:"event"`
|
||||||
|
} `json:"waiting"`
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(dir, "alerts.json")
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
|
||||||
|
if err == nil {
|
||||||
|
err = json.Unmarshal(data, &file)
|
||||||
|
}
|
||||||
|
|
||||||
|
waiting := file.Waiting["webhook"]
|
||||||
|
if err != nil || len(waiting) != 1 || waiting[0].Event != "permanent_ban" {
|
||||||
|
t.Fatalf("alerts.json holds %s (%v), want the permanent_ban alert waiting for "+
|
||||||
|
"the webhook", data, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// It counts the alert sent in the metrics, read here from a client the
|
||||||
|
// ban does not cover.
|
||||||
|
const token = "0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
webhook.failing.Store(false)
|
||||||
|
|
||||||
|
env["SWWAF_ALLOW_NETS"] = localhost
|
||||||
|
env[metricsToken] = token
|
||||||
|
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
webhook.waitFor(t, "permanent_ban", true)
|
||||||
|
|
||||||
|
const ofWebhook = `{destination="webhook",instance="fsn1app1/gitea"}`
|
||||||
|
|
||||||
|
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token,
|
||||||
|
"\nsmallwebwaf_alerts_sent_total"+ofWebhook+" 1\n")
|
||||||
|
|
||||||
|
for _, series := range []string{"failed", "suppressed", "dropped"} {
|
||||||
|
zero := "\nsmallwebwaf_alerts_" + series + "_total" + ofWebhook + " 0\n"
|
||||||
|
if !strings.Contains(metrics, zero) {
|
||||||
|
t.Errorf("no %q in the metrics:\n%s", zero, metrics)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBanIsAlertedToSlackAndNtfy(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
ntfyToken = "tk_0123456789abcdefghijklmnopq"
|
||||||
|
token = "abcdef0123456789abcdef0123456789"
|
||||||
|
client = "203.0.113.9"
|
||||||
|
)
|
||||||
|
|
||||||
|
slack, ntfy := startDestination(t), startDestination(t)
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
trustedProxies: localhost + "/32",
|
||||||
|
rateLimitPerDay: "1",
|
||||||
|
// The metrics are read from 127.0.0.1, which no limit counts.
|
||||||
|
"SWWAF_ALLOW_NETS": localhost + "/32",
|
||||||
|
metricsToken: token,
|
||||||
|
instanceName: instance,
|
||||||
|
"SWWAF_ALERT_SLACK_WEBHOOK_URL": slack.url,
|
||||||
|
"SWWAF_ALERT_NTFY_URL": ntfy.url,
|
||||||
|
"SWWAF_ALERT_NTFY_TOKEN": ntfyToken,
|
||||||
|
}
|
||||||
|
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
// The client's second request breaks the day limit, and bans it;
|
||||||
|
// Slack and ntfy are each sent the alert.
|
||||||
|
wantStatus(t, url, client, http.StatusOK)
|
||||||
|
wantStatus(t, url, client, http.StatusForbidden)
|
||||||
|
|
||||||
|
var message struct {
|
||||||
|
Text string `json:"text"`
|
||||||
|
}
|
||||||
|
|
||||||
|
slackPost := slack.firstPost(t)
|
||||||
|
err := json.Unmarshal([]byte(slackPost.body), &message)
|
||||||
|
|
||||||
|
if err != nil || !strings.HasPrefix(message.Text, "*fsn1app1/gitea: ban*\n") ||
|
||||||
|
!strings.Contains(message.Text, "\nclient: "+client+"\n") {
|
||||||
|
t.Errorf("Slack was sent %s", slackPost.body)
|
||||||
|
}
|
||||||
|
|
||||||
|
ntfyPost := ntfy.firstPost(t)
|
||||||
|
if ntfyPost.header.Get("Title") != "fsn1app1/gitea: ban" ||
|
||||||
|
ntfyPost.header.Get("Authorization") != "Bearer "+ntfyToken ||
|
||||||
|
!strings.Contains(ntfyPost.body, "\nclient: "+client+"\n") {
|
||||||
|
t.Errorf("ntfy was sent %s, with the headers %v", ntfyPost.body,
|
||||||
|
ntfyPost.header)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The metrics count it for each, and give no series for the
|
||||||
|
// webhook, which is not set. As long as that takes, so that a slow
|
||||||
|
// test process cannot fail the test.
|
||||||
|
sent := []string{
|
||||||
|
"\nsmallwebwaf_alerts_sent_total{destination=\"slack\"," +
|
||||||
|
"instance=\"fsn1app1/gitea\"} 1\n",
|
||||||
|
"\nsmallwebwaf_alerts_sent_total{destination=\"ntfy\"," +
|
||||||
|
"instance=\"fsn1app1/gitea\"} 1\n",
|
||||||
|
}
|
||||||
|
|
||||||
|
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token, sent...)
|
||||||
|
if strings.Contains(metrics, `destination="webhook"`) {
|
||||||
|
t.Errorf("the metrics give the webhook:\n%s", metrics)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
||||||
@@ -806,6 +1132,70 @@ func metricsText(t *testing.T, url, token string) string {
|
|||||||
return string(body)
|
return string(body)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// metricsWith asks for the metrics at url with token until they hold each
|
||||||
|
// of series, as long as that takes, so that a slow test process cannot
|
||||||
|
// fail the test, and returns them.
|
||||||
|
func metricsWith(t *testing.T, url, token string, series ...string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for {
|
||||||
|
metrics := metricsText(t, url, token)
|
||||||
|
|
||||||
|
missing := slices.ContainsFunc(series, func(one string) bool {
|
||||||
|
return !strings.Contains(metrics, one)
|
||||||
|
})
|
||||||
|
if !missing {
|
||||||
|
return metrics
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// seriesValue returns the value of series, such as
|
||||||
|
// name{instance="app"}, in metrics.
|
||||||
|
func seriesValue(t *testing.T, metrics, series string) float64 {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for line := range strings.Lines(metrics) {
|
||||||
|
value, found := strings.CutPrefix(strings.TrimSpace(line), series+" ")
|
||||||
|
if !found {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
number, err := strconv.ParseFloat(value, 64)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("%s has the value %q: %v", series, value, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return number
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Fatalf("no %s in the metrics:\n%s", series, metrics)
|
||||||
|
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// rename renames the file at from to, replacing any file there.
|
||||||
|
func rename(t *testing.T, from, to string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
err := os.Rename(from, to)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("rename %s to %s: %v", from, to, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeLookupDatabase writes a lookup database at path that places the
|
||||||
|
// client placed in country, and no other address.
|
||||||
|
func writeLookupDatabase(t *testing.T, path, country string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||||
|
placed + "/32": {ASN: "AS64496", ASName: "Example Net", Country: country},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// askAsAdmin sends a request with method to url, with body and
|
// askAsAdmin sends a request with method to url, with body and
|
||||||
// adminSecret, and checks that it is answered 200.
|
// adminSecret, and checks that it is answered 200.
|
||||||
func askAsAdmin(t *testing.T, method, url, body string) {
|
func askAsAdmin(t *testing.T, method, url, body string) {
|
||||||
@@ -917,6 +1307,140 @@ func saveUntilAnswered(t *testing.T, path, content, url, from string, status int
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// destination is a stand-in for SWWAF_ALERT_SLACK_WEBHOOK_URL or
|
||||||
|
// SWWAF_ALERT_NTFY_URL. It notes each request it is sent, and answers
|
||||||
|
// 200.
|
||||||
|
type destination struct {
|
||||||
|
url string
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
posts []destinationPost
|
||||||
|
}
|
||||||
|
|
||||||
|
// destinationPost is a request a destination was sent: its headers and
|
||||||
|
// its body.
|
||||||
|
type destinationPost struct {
|
||||||
|
header http.Header
|
||||||
|
body string
|
||||||
|
}
|
||||||
|
|
||||||
|
// startDestination starts a destination that takes every alert.
|
||||||
|
func startDestination(t *testing.T) *destination {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
d := &destination{}
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(
|
||||||
|
func(_ http.ResponseWriter, r *http.Request) {
|
||||||
|
body, _ := io.ReadAll(r.Body)
|
||||||
|
|
||||||
|
d.mu.Lock()
|
||||||
|
d.posts = append(d.posts, destinationPost{
|
||||||
|
header: r.Header.Clone(), body: string(body),
|
||||||
|
})
|
||||||
|
d.mu.Unlock()
|
||||||
|
}))
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
d.url = server.URL + "/alerts"
|
||||||
|
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstPost waits until the destination has been sent a request, and
|
||||||
|
// returns the first. It waits as long as that takes, so that a slow test
|
||||||
|
// process cannot fail the test.
|
||||||
|
func (d *destination) firstPost(t *testing.T) destinationPost {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for {
|
||||||
|
d.mu.Lock()
|
||||||
|
|
||||||
|
if len(d.posts) > 0 {
|
||||||
|
post := d.posts[0]
|
||||||
|
d.mu.Unlock()
|
||||||
|
|
||||||
|
return post
|
||||||
|
}
|
||||||
|
|
||||||
|
d.mu.Unlock()
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// webhook is a stand-in for SWWAF_ALERT_WEBHOOK_URL. It notes each alert
|
||||||
|
// it is sent, and answers 204, or 503 while failing.
|
||||||
|
type webhook struct {
|
||||||
|
url string
|
||||||
|
failing atomic.Bool
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
posts []webhookPost
|
||||||
|
}
|
||||||
|
|
||||||
|
// webhookPost is an alert the webhook was sent, with the Authorization
|
||||||
|
// header sent with it, and whether the webhook took it.
|
||||||
|
type webhookPost struct {
|
||||||
|
alert map[string]any
|
||||||
|
authorization string
|
||||||
|
answered bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// startWebhook starts a webhook that takes every alert.
|
||||||
|
func startWebhook(t *testing.T) *webhook {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
w := &webhook{}
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(
|
||||||
|
func(rw http.ResponseWriter, r *http.Request) {
|
||||||
|
var alert map[string]any
|
||||||
|
|
||||||
|
_ = json.NewDecoder(r.Body).Decode(&alert)
|
||||||
|
failing := w.failing.Load()
|
||||||
|
|
||||||
|
w.mu.Lock()
|
||||||
|
w.posts = append(w.posts, webhookPost{
|
||||||
|
alert: alert, authorization: r.Header.Get("Authorization"),
|
||||||
|
answered: !failing,
|
||||||
|
})
|
||||||
|
w.mu.Unlock()
|
||||||
|
|
||||||
|
if failing {
|
||||||
|
rw.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
rw.WriteHeader(http.StatusNoContent)
|
||||||
|
}))
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
w.url = server.URL + "/alerts"
|
||||||
|
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitFor waits until the webhook has been sent an alert for event that
|
||||||
|
// it took, or, unless answered, failed, and returns it. It waits as long
|
||||||
|
// as that takes, so that a slow test process cannot fail the test.
|
||||||
|
func (w *webhook) waitFor(t *testing.T, event string, answered bool) webhookPost {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for {
|
||||||
|
w.mu.Lock()
|
||||||
|
|
||||||
|
for _, post := range w.posts {
|
||||||
|
if post.alert["event"] == event && post.answered == answered {
|
||||||
|
w.mu.Unlock()
|
||||||
|
|
||||||
|
return post
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
w.mu.Unlock()
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// statusFrom returns the status a request to url from the client at
|
// statusFrom returns the status a request to url from the client at
|
||||||
// from, as X-Forwarded-For names it, is answered with.
|
// from, as X-Forwarded-For names it, is answered with.
|
||||||
func statusFrom(t *testing.T, url, from string) int {
|
func statusFrom(t *testing.T, url, from string) int {
|
||||||
|
|||||||
+152
-30
@@ -1,11 +1,13 @@
|
|||||||
// Package state keeps smallwebwaf's state in JSON files in
|
// Package state keeps smallwebwaf's state in JSON files in
|
||||||
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
|
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
|
||||||
// bans.json holds the bans, clients.json each client's counters and
|
// bans.json holds the bans, clients.json each client's counters and
|
||||||
// history, and lookups.json GeoJS's answers. Load reads them at start,
|
// history, lookups.json GeoJS's answers, and alerts.json the cooldowns,
|
||||||
// Watch takes in an admin's edit of one while smallwebwaf runs, and Run
|
// the hour under way and the alerts waiting for each destination. Load
|
||||||
// and WriteAll write them. The disk is read and written outside the
|
// reads them at start, Watch takes in an admin's edit of one while
|
||||||
// parts' locks, which are held only to take a snapshot or to put in what
|
// smallwebwaf runs, and Run and WriteAll write them. The disk is read and
|
||||||
// a file holds, so that no request waits on the disk.
|
// written outside the parts' locks, which are held only to take a
|
||||||
|
// snapshot or to put in what a file holds, so that no request waits on
|
||||||
|
// the disk.
|
||||||
package state
|
package state
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -17,13 +19,16 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"maps"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
"github.com/fsnotify/fsnotify"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
@@ -42,13 +47,18 @@ const (
|
|||||||
bansJSON = "bans.json"
|
bansJSON = "bans.json"
|
||||||
clientsJSON = "clients.json"
|
clientsJSON = "clients.json"
|
||||||
lookupsJSON = "lookups.json"
|
lookupsJSON = "lookups.json"
|
||||||
|
alertsJSON = "alerts.json"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
errVersion = errors.New("unknown version")
|
errVersion = errors.New("unknown version")
|
||||||
// errMissing is for an entry without a field it needs.
|
// errMissing is for an entry without a field it needs.
|
||||||
errMissing = errors.New("has no")
|
errMissing = errors.New("has no")
|
||||||
errCause = errors.New("is not limit, attack or admin")
|
errCause = errors.New("is not limit, attack or admin")
|
||||||
|
errDestination = errors.New("is not webhook, slack or ntfy")
|
||||||
|
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
|
||||||
|
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
|
||||||
|
`or remove the file`)
|
||||||
)
|
)
|
||||||
|
|
||||||
// Params are what Load needs.
|
// Params are what Load needs.
|
||||||
@@ -60,10 +70,13 @@ type Params struct {
|
|||||||
// is (SWWAF_STATE_COUNTER_INTERVAL).
|
// is (SWWAF_STATE_COUNTER_INTERVAL).
|
||||||
WriteDelay time.Duration
|
WriteDelay time.Duration
|
||||||
CounterInterval time.Duration
|
CounterInterval time.Duration
|
||||||
// Ledger, Limiter and GeoJS hold the state.
|
// Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
|
||||||
|
// receive a file_error alert for an edit set aside, and for a write
|
||||||
|
// that fails while smallwebwaf runs.
|
||||||
Ledger *bans.Ledger
|
Ledger *bans.Ledger
|
||||||
Limiter *ratelimit.Limiter
|
Limiter *ratelimit.Limiter
|
||||||
GeoJS *lookup.GeoJS
|
GeoJS *lookup.GeoJS
|
||||||
|
Alerts *alerts.Queue
|
||||||
// Now tells the time by which the counters' buckets run out, normally
|
// Now tells the time by which the counters' buckets run out, normally
|
||||||
// time.Now in UTC.
|
// time.Now in UTC.
|
||||||
Now func() time.Time
|
Now func() time.Time
|
||||||
@@ -121,6 +134,14 @@ type lookupsFile struct {
|
|||||||
Lookups []lookup.Answer `json:"lookups"`
|
Lookups []lookup.Answer `json:"lookups"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// alertsFile is alerts.json, indented for an admin to read and edit.
|
||||||
|
type alertsFile struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Cooldowns []alerts.Cooldown `json:"cooldowns"`
|
||||||
|
Hour alerts.Hour `json:"hour"`
|
||||||
|
Waiting map[string][]alerts.Alert `json:"waiting"`
|
||||||
|
}
|
||||||
|
|
||||||
// stateFile is the struct of a state file. Once the file is decoded, its
|
// stateFile is the struct of a state file. Once the file is decoded, its
|
||||||
// check refuses the first entry without a field it needs, which would
|
// check refuses the first entry without a field it needs, which would
|
||||||
// otherwise be read as something the entry does not say. data is the
|
// otherwise be read as something the entry does not say. data is the
|
||||||
@@ -147,23 +168,25 @@ func Load(params Params) (*Files, error) {
|
|||||||
bansRead, bansErr := f.read(bansJSON)
|
bansRead, bansErr := f.read(bansJSON)
|
||||||
clientsRead, clientsErr := f.read(clientsJSON)
|
clientsRead, clientsErr := f.read(clientsJSON)
|
||||||
lookupsRead, lookupsErr := f.read(lookupsJSON)
|
lookupsRead, lookupsErr := f.read(lookupsJSON)
|
||||||
|
alertsRead, alertsErr := f.read(alertsJSON)
|
||||||
|
|
||||||
err = errors.Join(bansErr, clientsErr, lookupsErr)
|
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
||||||
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
|
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
|
||||||
|
"alerts_waiting", alertsRead)
|
||||||
|
|
||||||
return f, nil
|
return f, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run writes bans.json WriteDelay after a ban is made, with every ban
|
// Run writes bans.json WriteDelay after a ban is made, with every ban
|
||||||
// made in between, and every file every CounterInterval, until ctx is
|
// made in between, and every file every CounterInterval, until ctx is
|
||||||
// done. A write that fails is logged, and the file is written again at
|
// done. A write that fails is logged, raised as a file_error alert, and
|
||||||
// its next write. Each write takes in an admin's edit of its file first,
|
// the file is written again at its next write. Each write takes in an
|
||||||
// as writeFile describes.
|
// admin's edit of its file first, as writeFile describes.
|
||||||
func (f *Files) Run(ctx context.Context) {
|
func (f *Files) Run(ctx context.Context) {
|
||||||
interval := time.NewTicker(f.params.CounterInterval)
|
interval := time.NewTicker(f.params.CounterInterval)
|
||||||
defer interval.Stop()
|
defer interval.Stop()
|
||||||
@@ -181,9 +204,11 @@ func (f *Files) Run(ctx context.Context) {
|
|||||||
case <-bansDue:
|
case <-bansDue:
|
||||||
bansDue = nil
|
bansDue = nil
|
||||||
|
|
||||||
f.logFailure(f.writeFile(bansJSON))
|
f.logFailure(bansJSON, f.writeFile(bansJSON))
|
||||||
case <-interval.C:
|
case <-interval.C:
|
||||||
f.logFailure(f.WriteAll())
|
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
|
||||||
|
f.logFailure(name, f.writeFile(name))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -192,7 +217,7 @@ func (f *Files) Run(ctx context.Context) {
|
|||||||
// fails does not keep the others from being written.
|
// fails does not keep the others from being written.
|
||||||
func (f *Files) WriteAll() error {
|
func (f *Files) WriteAll() error {
|
||||||
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
|
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
|
||||||
f.writeFile(lookupsJSON))
|
f.writeFile(lookupsJSON), f.writeFile(alertsJSON))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
|
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
|
||||||
@@ -227,7 +252,7 @@ func (f *Files) Watch(ctx context.Context) {
|
|||||||
return
|
return
|
||||||
case event := <-watcher.Events:
|
case event := <-watcher.Events:
|
||||||
switch name := filepath.Base(event.Name); name {
|
switch name := filepath.Base(event.Name); name {
|
||||||
case bansJSON, clientsJSON, lookupsJSON:
|
case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
|
||||||
f.fileChanged(name)
|
f.fileChanged(name)
|
||||||
}
|
}
|
||||||
case err = <-watcher.Errors:
|
case err = <-watcher.Errors:
|
||||||
@@ -237,11 +262,22 @@ func (f *Files) Watch(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// logFailure logs a write that failed.
|
// logFailure logs a write of the state file name that failed, and raises
|
||||||
func (f *Files) logFailure(err error) {
|
// a file_error alert for it.
|
||||||
|
func (f *Files) logFailure(name string, err error) {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.params.ProcessLog.Error("writing the state files failed",
|
const failed = "writing the state files failed"
|
||||||
"error", err.Error())
|
|
||||||
|
// Raised before it is logged, so that the alert is there once the
|
||||||
|
// log line is.
|
||||||
|
f.params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventFileError,
|
||||||
|
Reason: failed,
|
||||||
|
Detail: map[string]any{
|
||||||
|
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
f.params.ProcessLog.Error(failed, "error", err.Error())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -362,6 +398,32 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
|
|||||||
|
|
||||||
f.params.GeoJS.Load(file.Lookups)
|
f.params.GeoJS.Load(file.Lookups)
|
||||||
entries = len(file.Lookups)
|
entries = len(file.Lookups)
|
||||||
|
case alertsJSON:
|
||||||
|
// waiting was a list, of the alerts waiting for the webhook, before
|
||||||
|
// alerts went to Slack and ntfy too.
|
||||||
|
var written struct {
|
||||||
|
Waiting json.RawMessage `json:"waiting"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if json.Unmarshal(data, &written) == nil &&
|
||||||
|
bytes.HasPrefix(written.Waiting, []byte("[")) {
|
||||||
|
return 0, fmt.Errorf("%s: %w", path, errWaitingList)
|
||||||
|
}
|
||||||
|
|
||||||
|
var file alertsFile
|
||||||
|
|
||||||
|
err := parse(path, data, &file)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.Alerts.Load(alerts.State{
|
||||||
|
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, waiting := range file.Waiting {
|
||||||
|
entries += len(waiting)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.sums[name] = sha256.Sum256(data)
|
f.sums[name] = sha256.Sum256(data)
|
||||||
@@ -413,8 +475,9 @@ func (f *Files) writeFile(name string) error {
|
|||||||
|
|
||||||
// setAside renames the state file name, an edit that does not parse with
|
// setAside renames the state file name, an edit that does not parse with
|
||||||
// parseErr, to name.bad, for the admin to mend, and logs it with where in
|
// parseErr, to name.bad, for the admin to mend, and logs it with where in
|
||||||
// the file the error is. If the rename fails, the edit is left as it is,
|
// the file the error is, and raises a file_error alert for it. If the
|
||||||
// and the error returned is parseErr joined with the rename's.
|
// rename fails, the edit is left as it is, and the error returned is
|
||||||
|
// parseErr joined with the rename's.
|
||||||
func (f *Files) setAside(name string, parseErr error) error {
|
func (f *Files) setAside(name string, parseErr error) error {
|
||||||
path := filepath.Join(f.params.Dir, name)
|
path := filepath.Join(f.params.Dir, name)
|
||||||
|
|
||||||
@@ -423,8 +486,16 @@ func (f *Files) setAside(name string, parseErr error) error {
|
|||||||
return errors.Join(parseErr, err)
|
return errors.Join(parseErr, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse",
|
const setAside = "set aside an edit of a state file that does not parse"
|
||||||
"file", path+".bad", "error", parseErr.Error())
|
|
||||||
|
// Raised before it is logged, so that the alert is there once the log
|
||||||
|
// line is.
|
||||||
|
f.params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventFileError,
|
||||||
|
Reason: setAside,
|
||||||
|
Detail: map[string]any{"file": path + ".bad", "error": parseErr.Error()},
|
||||||
|
})
|
||||||
|
f.params.ProcessLog.Error(setAside, "file", path+".bad", "error", parseErr.Error())
|
||||||
f.params.Metrics.StateFileEditSetAside(name)
|
f.params.Metrics.StateFileEditSetAside(name)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -445,8 +516,21 @@ func (f *Files) encode(name string) ([]byte, error) {
|
|||||||
return append(data, '\n'), nil
|
return append(data, '\n'), nil
|
||||||
case clientsJSON:
|
case clientsJSON:
|
||||||
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
||||||
default: // lookups.json
|
case lookupsJSON:
|
||||||
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
||||||
|
default: // alerts.json
|
||||||
|
held := f.params.Alerts.Snapshot()
|
||||||
|
file := alertsFile{
|
||||||
|
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
|
||||||
|
Waiting: held.Waiting,
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.MarshalIndent(file, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return append(data, '\n'), nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -531,8 +615,8 @@ func (f *bansFile) check(data []byte) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// check refuses a client without its address, which would count nobody's
|
// check refuses a client without its address, which would count nobody's
|
||||||
// requests, or with requests in a window but no start, which would drop
|
// requests, or with requests or bytes in a window but no start, which
|
||||||
// them and give the client a fresh allowance.
|
// would drop them and give the client a fresh allowance.
|
||||||
func (f *clientsFile) check([]byte) error {
|
func (f *clientsFile) check([]byte) error {
|
||||||
for i, client := range f.Clients {
|
for i, client := range f.Clients {
|
||||||
switch {
|
switch {
|
||||||
@@ -544,6 +628,12 @@ func (f *clientsFile) check([]byte) error {
|
|||||||
return missing(i, "hour.start")
|
return missing(i, "hour.start")
|
||||||
case countsWithoutStart(client.Day):
|
case countsWithoutStart(client.Day):
|
||||||
return missing(i, "day.start")
|
return missing(i, "day.start")
|
||||||
|
case countsWithoutStart(client.MinuteBytes):
|
||||||
|
return missing(i, "minute_bytes.start")
|
||||||
|
case countsWithoutStart(client.HourBytes):
|
||||||
|
return missing(i, "hour_bytes.start")
|
||||||
|
case countsWithoutStart(client.DayBytes):
|
||||||
|
return missing(i, "day_bytes.start")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -581,8 +671,40 @@ func (f *lookupsFile) check(data []byte) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// countsWithoutStart reports whether b holds requests but no start, which
|
// check refuses a cooldown without its event or when its alert was sent,
|
||||||
// places them in time.
|
// which would hold back no repeat, alerts waiting for a destination with
|
||||||
|
// another name than webhook, slack or ntfy, most likely misspelt, and an
|
||||||
|
// alert waiting without its event or its time.
|
||||||
|
func (f *alertsFile) check([]byte) error {
|
||||||
|
for i, cooldown := range f.Cooldowns {
|
||||||
|
switch {
|
||||||
|
case cooldown.Event == "":
|
||||||
|
return fmt.Errorf("cooldowns %w", missing(i, "event"))
|
||||||
|
case cooldown.Sent.IsZero():
|
||||||
|
return fmt.Errorf("cooldowns %w", missing(i, "sent"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, destination := range slices.Sorted(maps.Keys(f.Waiting)) {
|
||||||
|
if !slices.Contains(alerts.Destinations(), destination) {
|
||||||
|
return fmt.Errorf("waiting %q %w", destination, errDestination)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, alert := range f.Waiting[destination] {
|
||||||
|
switch {
|
||||||
|
case alert.Event == "":
|
||||||
|
return fmt.Errorf("waiting %s %w", destination, missing(i, "event"))
|
||||||
|
case alert.Time.IsZero():
|
||||||
|
return fmt.Errorf("waiting %s %w", destination, missing(i, "time"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// countsWithoutStart reports whether b holds requests, or bytes, but no
|
||||||
|
// start, which places them in time.
|
||||||
func countsWithoutStart(b ratelimit.Buckets) bool {
|
func countsWithoutStart(b ratelimit.Buckets) bool {
|
||||||
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
|
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
|
||||||
}
|
}
|
||||||
|
|||||||
+352
-32
@@ -10,8 +10,10 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -19,6 +21,7 @@ import (
|
|||||||
"testing/synctest"
|
"testing/synctest"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
@@ -31,6 +34,10 @@ const (
|
|||||||
bansJSON = "bans.json"
|
bansJSON = "bans.json"
|
||||||
clientsJSON = "clients.json"
|
clientsJSON = "clients.json"
|
||||||
lookupsJSON = "lookups.json"
|
lookupsJSON = "lookups.json"
|
||||||
|
alertsJSON = "alerts.json"
|
||||||
|
// The AS number and AS name the tests' clients are looked up in.
|
||||||
|
asn = "AS64496"
|
||||||
|
asName = "Example Net"
|
||||||
// What the process log says once Watch watches the directory, and as
|
// What the process log says once Watch watches the directory, and as
|
||||||
// it takes in an edit.
|
// it takes in an edit.
|
||||||
watching = "watching the state files for edits"
|
watching = "watching the state files for edits"
|
||||||
@@ -38,6 +45,9 @@ const (
|
|||||||
// maxLogLines is how many lines of the process log wait for a test to
|
// maxLogLines is how many lines of the process log wait for a test to
|
||||||
// read them.
|
// read them.
|
||||||
maxLogLines = 64
|
maxLogLines = 64
|
||||||
|
// whole is the percentage of each limit a client gets when nothing
|
||||||
|
// lowers its limits.
|
||||||
|
whole = 100
|
||||||
)
|
)
|
||||||
|
|
||||||
// permanentBansJSON is bans.json holding permanentBan.
|
// permanentBansJSON is bans.json holding permanentBan.
|
||||||
@@ -51,6 +61,8 @@ const permanentBansJSON = `{
|
|||||||
"cause": "admin",
|
"cause": "admin",
|
||||||
"reason": "scrapes every commit",
|
"reason": "scrapes every commit",
|
||||||
"notes": {
|
"notes": {
|
||||||
|
"asn": "AS64496",
|
||||||
|
"as_name": "Example Net",
|
||||||
"country": "DE",
|
"country": "DE",
|
||||||
"limit": 1000,
|
"limit": 1000,
|
||||||
"window": "minute",
|
"window": "minute",
|
||||||
@@ -85,6 +97,69 @@ const liftedBansJSON = `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", `
|
|||||||
`"start": "2026-10-06T00:00:00Z", "expires": "2026-10-06T01:00:00Z", ` +
|
`"start": "2026-10-06T00:00:00Z", "expires": "2026-10-06T01:00:00Z", ` +
|
||||||
`"cause": "limit", "lifted": "2026-10-06T00:10:00Z"}]}`
|
`"cause": "limit", "lifted": "2026-10-06T00:10:00Z"}]}`
|
||||||
|
|
||||||
|
// filledAlertsJSON is alerts.json holding the alerts of fill.
|
||||||
|
const filledAlertsJSON = `{
|
||||||
|
"version": 1,
|
||||||
|
"cooldowns": [
|
||||||
|
{
|
||||||
|
"event": "file_error",
|
||||||
|
"netblock": "",
|
||||||
|
"file": "/var/lib/smallwebwaf/bans.json",
|
||||||
|
"sent": "2026-10-06T00:00:00Z",
|
||||||
|
"suppressed_repeats": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"event": "ban",
|
||||||
|
"netblock": "203.0.113.9/32",
|
||||||
|
"sent": "2026-10-06T00:00:00Z",
|
||||||
|
"suppressed_repeats": 1
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"hour": {
|
||||||
|
"start": "2026-10-06T00:00:00Z",
|
||||||
|
"sent": 2,
|
||||||
|
"held_back": {
|
||||||
|
"source_failure": 1
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"waiting": {
|
||||||
|
"webhook": [
|
||||||
|
{
|
||||||
|
"instance": "fsn1app1/gitea",
|
||||||
|
"time": "2026-10-06T00:00:00Z",
|
||||||
|
"event": "ban",
|
||||||
|
"client": "203.0.113.9",
|
||||||
|
"netblock": "203.0.113.9/32",
|
||||||
|
"asn": "",
|
||||||
|
"as_name": "",
|
||||||
|
"country": "DE",
|
||||||
|
"reason": "requests per minute over the limit of 1",
|
||||||
|
"detail": {
|
||||||
|
"cause": "limit"
|
||||||
|
},
|
||||||
|
"suppressed_repeats": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"instance": "fsn1app1/gitea",
|
||||||
|
"time": "2026-10-06T00:00:00Z",
|
||||||
|
"event": "file_error",
|
||||||
|
"client": "",
|
||||||
|
"netblock": "",
|
||||||
|
"asn": "",
|
||||||
|
"as_name": "",
|
||||||
|
"country": "",
|
||||||
|
"reason": "writing the state files failed",
|
||||||
|
"detail": {
|
||||||
|
"error": "no space left on device",
|
||||||
|
"file": "/var/lib/smallwebwaf/bans.json"
|
||||||
|
},
|
||||||
|
"suppressed_repeats": 0
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
`
|
||||||
|
|
||||||
func TestFilesWrittenAndReadBack(t *testing.T) {
|
func TestFilesWrittenAndReadBack(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -111,13 +186,69 @@ func TestFilesWrittenAndReadBack(t *testing.T) {
|
|||||||
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
|
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
|
||||||
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
|
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
|
||||||
|
|
||||||
|
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
|
||||||
|
got, want) {
|
||||||
|
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
|
||||||
|
}
|
||||||
|
|
||||||
// Each one-per-line file lists its entries by client, and nothing
|
// Each one-per-line file lists its entries by client, and nothing
|
||||||
// but the three files is left in the directory.
|
// but the four files is left in the directory.
|
||||||
wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
|
wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
|
||||||
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
|
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
|
||||||
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
|
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
|
||||||
"192.0.2.1/32", "203.0.113.9/32")
|
"192.0.2.1/32", "203.0.113.9/32")
|
||||||
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
params := newParams(dir)
|
||||||
|
fill(params)
|
||||||
|
|
||||||
|
files := load(t, params)
|
||||||
|
|
||||||
|
err := files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := readFile(t, filepath.Join(dir, alertsJSON))
|
||||||
|
if got != filledAlertsJSON {
|
||||||
|
t.Errorf("alerts.json\n%s\nwant\n%s", got, filledAlertsJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSourceFailureCooldownKeptInAlertsJSONAcrossARestart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
before := newParams(dir)
|
||||||
|
files := load(t, before)
|
||||||
|
|
||||||
|
failure := alerts.Alert{
|
||||||
|
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
|
||||||
|
Detail: map[string]any{"source": "geojs"},
|
||||||
|
}
|
||||||
|
before.Alerts.Raise(failure)
|
||||||
|
|
||||||
|
err := files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// After the restart, the cooldown read back holds back a repeat for the
|
||||||
|
// same source.
|
||||||
|
after := newParams(dir)
|
||||||
|
load(t, after)
|
||||||
|
after.Alerts.Raise(failure)
|
||||||
|
|
||||||
|
waiting := after.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 || after.Alerts.Suppressed() != 1 {
|
||||||
|
t.Errorf("%d alerts wait and %d are held back, want the one read back and 1",
|
||||||
|
len(waiting), after.Alerts.Suppressed())
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
|
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
|
||||||
@@ -146,8 +277,10 @@ func TestMissingFilesAreEmptyState(t *testing.T) {
|
|||||||
params := newParams(t.TempDir())
|
params := newParams(t.TempDir())
|
||||||
load(t, params)
|
load(t, params)
|
||||||
|
|
||||||
|
held := params.Alerts.Snapshot()
|
||||||
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
|
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
|
||||||
len(params.GeoJS.Snapshot()) != 0 {
|
len(params.GeoJS.Snapshot()) != 0 || len(held.Cooldowns) != 0 ||
|
||||||
|
len(held.Waiting[alerts.DestinationWebhook]) != 0 || held.Hour.Sent != 0 {
|
||||||
t.Error("state from no files")
|
t.Error("state from no files")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -189,6 +322,16 @@ func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
|||||||
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
|
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
|
||||||
`: netip.ParsePrefix("203.0.113.300/32")`,
|
`: netip.ParsePrefix("203.0.113.300/32")`,
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"an unknown field of an alert waiting", alertsJSON,
|
||||||
|
`{"version": 1, "waiting": {"webhook": [{"event": "ban", "evnet": "ban"}]}}`,
|
||||||
|
`: json: unknown field "evnet"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"alerts waiting for an unknown destination", alertsJSON,
|
||||||
|
`{"version": 1, "waiting": {"webhook": [], "slak": []}}`,
|
||||||
|
`: waiting "slak" is not webhook, slack or ntfy`,
|
||||||
|
},
|
||||||
} {
|
} {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -252,6 +395,12 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
|||||||
`"hour": {"current": 3}}]}`,
|
`"hour": {"current": 3}}]}`,
|
||||||
`: entry 1 has no "hour.start"`,
|
`: entry 1 has no "hour.start"`,
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"a client with bytes in a window without its start", clientsJSON,
|
||||||
|
`{"version": 1, "clients": [{"client": "203.0.113.9/32", ` +
|
||||||
|
`"minute_bytes": {"previous": 5120}}]}`,
|
||||||
|
`: entry 1 has no "minute_bytes.start"`,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"an answer without a client", lookupsJSON,
|
"an answer without a client", lookupsJSON,
|
||||||
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
|
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
|
||||||
@@ -278,6 +427,60 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, content string
|
||||||
|
// want is what the error says after the file's path.
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"a cooldown without its event",
|
||||||
|
`{"version": 1, "cooldowns": [{"sent": "2026-10-06T00:00:00Z"}]}`,
|
||||||
|
`: cooldowns entry 1 has no "event"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a cooldown without when it was sent",
|
||||||
|
`{"version": 1, "cooldowns": [{"event": "ban"}]}`,
|
||||||
|
`: cooldowns entry 1 has no "sent"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an alert waiting without its event",
|
||||||
|
`{"version": 1, "waiting": {"ntfy": [{"time": "2026-10-06T00:00:00Z"}]}}`,
|
||||||
|
`: waiting ntfy entry 1 has no "event"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an alert waiting without its time",
|
||||||
|
`{"version": 1, "waiting": {"slack": [{"event": "ban"}]}}`,
|
||||||
|
`: waiting slack entry 1 has no "time"`,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantRefused(t, alertsJSON, tc.content, tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAlertsJSONWithWaitingAsAListStopsTheStartSayingWhatToChange(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// alerts.json as it was written before alerts went to Slack and ntfy too,
|
||||||
|
// with no alert waiting, or one.
|
||||||
|
for _, waiting := range []string{
|
||||||
|
`[]`,
|
||||||
|
`[{"event": "ban", "time": "2026-10-06T00:00:00Z"}]`,
|
||||||
|
} {
|
||||||
|
wantRefused(t, alertsJSON, `{"version": 1, "cooldowns": [], `+
|
||||||
|
`"hour": {"start": "2026-10-06T00:00:00Z", "sent": 0, "held_back": {}}, `+
|
||||||
|
`"waiting": `+waiting+`}`,
|
||||||
|
`: waiting is a list, but now lists the alerts by destination: put the `+
|
||||||
|
`list under "webhook", as "waiting": {"webhook": [...]}, or remove the file`)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
|
func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -294,7 +497,7 @@ func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
|
|||||||
func TestUnknownVersionStopsTheStart(t *testing.T) {
|
func TestUnknownVersionStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} {
|
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
|
||||||
for _, content := range []string{`{"version": 2}`, `{}`} {
|
for _, content := range []string{`{"version": 2}`, `{}`} {
|
||||||
t.Run(file+" "+content, func(t *testing.T) {
|
t.Run(file+" "+content, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -344,12 +547,12 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
|
|||||||
|
|
||||||
// A second ban, made while the first waits to be written, puts the
|
// A second ban, made while the first waits to be written, puts the
|
||||||
// write off no further, and is written with it.
|
// write off no further, and is written with it.
|
||||||
first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
|
first, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
|
||||||
midnight(), bans.Notes{})
|
midnight(), bans.Notes{})
|
||||||
|
|
||||||
time.Sleep(5 * time.Second)
|
time.Sleep(5 * time.Second)
|
||||||
|
|
||||||
second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
|
second, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
|
||||||
midnight(), bans.Notes{})
|
midnight(), bans.Notes{})
|
||||||
|
|
||||||
time.Sleep(5*time.Second - time.Nanosecond)
|
time.Sleep(5*time.Second - time.Nanosecond)
|
||||||
@@ -396,8 +599,8 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
|
|||||||
|
|
||||||
time.Sleep(time.Nanosecond)
|
time.Sleep(time.Nanosecond)
|
||||||
synctest.Wait()
|
synctest.Wait()
|
||||||
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||||
removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -440,6 +643,57 @@ func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestWriteThatFailsWhileRunningRaisesAFileErrorAlertOncePerCooldown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
params := newParams(dir)
|
||||||
|
params.CounterInterval = time.Minute
|
||||||
|
run(t, load(t, params).Run)
|
||||||
|
|
||||||
|
// A directory in the way of clients.json's temporary file fails each
|
||||||
|
// of its writes, while the other files are written. It holds a file,
|
||||||
|
// so that the write cannot remove it.
|
||||||
|
err := os.Mkdir(filepath.Join(dir, clientsJSON+".tmp"), 0o700)
|
||||||
|
if err == nil {
|
||||||
|
err = os.WriteFile(filepath.Join(dir, clientsJSON+".tmp", "kept"), nil, 0o600)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("put a directory in the way: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(time.Minute)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
waiting := params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 {
|
||||||
|
t.Fatalf("%d alerts wait, want 1", len(waiting))
|
||||||
|
}
|
||||||
|
|
||||||
|
message, _ := waiting[0].Detail["error"].(string)
|
||||||
|
|
||||||
|
if waiting[0].Event != alerts.EventFileError ||
|
||||||
|
waiting[0].Reason != "writing the state files failed" ||
|
||||||
|
waiting[0].Detail["file"] != filepath.Join(dir, clientsJSON) ||
|
||||||
|
!strings.Contains(message, clientsJSON+".tmp") {
|
||||||
|
t.Fatalf("alerts waiting %+v, want a file_error alert for clients.json, "+
|
||||||
|
"naming its temporary file", waiting)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next write fails too, within the cooldown, which holds it back.
|
||||||
|
time.Sleep(time.Minute)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
waiting = params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 || params.Alerts.Suppressed() != 1 {
|
||||||
|
t.Errorf("%d alerts wait and %d are held back, want 1 and 1",
|
||||||
|
len(waiting), params.Alerts.Suppressed())
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
|
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -462,7 +716,7 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
|
|||||||
|
|
||||||
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
||||||
bans.Notes{})
|
bans.Notes{})
|
||||||
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
|
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight(), whole)
|
||||||
|
|
||||||
err = files.WriteAll()
|
err = files.WriteAll()
|
||||||
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
|
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
|
||||||
@@ -496,8 +750,8 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
ofBans = `{file="bans.json"}`
|
ofBans = `{file="bans.json",instance="app"}`
|
||||||
ofClients = `{file="clients.json"}`
|
ofClients = `{file="clients.json",instance="app"}`
|
||||||
)
|
)
|
||||||
|
|
||||||
got := scrape(t, params)
|
got := scrape(t, params)
|
||||||
@@ -565,7 +819,7 @@ func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
|
|||||||
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
|
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||||
wantWriteFailed(t, params, bansJSON)
|
wantWriteFailed(t, params, bansJSON)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -637,6 +891,27 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
|
|||||||
wantTakenIn(t, lines, dir, lookupsJSON)
|
wantTakenIn(t, lines, dir, lookupsJSON)
|
||||||
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
|
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
|
||||||
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
|
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
|
||||||
|
|
||||||
|
// A netblock with bits past its length is read as the netblock it is
|
||||||
|
// in.
|
||||||
|
edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+
|
||||||
|
`"netblock": "198.51.100.9/24", "sent": "2026-10-06T00:00:00Z"}], `+
|
||||||
|
`"waiting": {"webhook": [{"event": "file_error", "time": "2026-10-06T00:00:00Z"}]}}`)
|
||||||
|
wantTakenIn(t, lines, dir, alertsJSON)
|
||||||
|
|
||||||
|
want := alerts.State{
|
||||||
|
Cooldowns: []alerts.Cooldown{{
|
||||||
|
Event: alerts.EventBan, Netblock: netip.MustParsePrefix("198.51.100.0/24"),
|
||||||
|
Sent: midnight(),
|
||||||
|
}},
|
||||||
|
Hour: alerts.Hour{HeldBack: map[string]int{}},
|
||||||
|
Waiting: map[string][]alerts.Alert{
|
||||||
|
alerts.DestinationWebhook: {{Event: alerts.EventFileError, Time: midnight()}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if got := params.Alerts.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("%s taken in as\n%+v\nwant\n%+v", alertsJSON, got, want)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestOwnWritesAreNotTakenIn(t *testing.T) {
|
func TestOwnWritesAreNotTakenIn(t *testing.T) {
|
||||||
@@ -708,7 +983,7 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
|
|||||||
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
|
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
|
||||||
wantTakenIn(t, lines, dir, bansJSON)
|
wantTakenIn(t, lines, dir, bansJSON)
|
||||||
|
|
||||||
_, banned := params.Ledger.Check(client, midnight())
|
_, banned, _ := params.Ledger.Check(client, midnight())
|
||||||
if !banned {
|
if !banned {
|
||||||
t.Error("the ban added to bans.json does not refuse")
|
t.Error("the ban added to bans.json does not refuse")
|
||||||
}
|
}
|
||||||
@@ -717,7 +992,7 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
|
|||||||
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
|
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
|
||||||
wantTakenIn(t, lines, dir, bansJSON)
|
wantTakenIn(t, lines, dir, bansJSON)
|
||||||
|
|
||||||
_, banned = params.Ledger.Check(client, midnight())
|
_, banned, _ = params.Ledger.Check(client, midnight())
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the ban removed from bans.json still refuses")
|
t.Error("the ban removed from bans.json still refuses")
|
||||||
}
|
}
|
||||||
@@ -807,7 +1082,7 @@ func TestBanLiftedByAnEditWhileRunning(t *testing.T) {
|
|||||||
netblock := netip.MustParsePrefix(liftedClient + "/32")
|
netblock := netip.MustParsePrefix(liftedClient + "/32")
|
||||||
params.Ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
params.Ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
|
|
||||||
_, banned := params.Ledger.Find(netblock.Addr(), afterLifting())
|
_, banned, _ := params.Ledger.Find(netblock.Addr(), afterLifting())
|
||||||
if !banned {
|
if !banned {
|
||||||
t.Fatal("the ban does not refuse before it is lifted")
|
t.Fatal("the ban does not refuse before it is lifted")
|
||||||
}
|
}
|
||||||
@@ -844,7 +1119,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
|
|||||||
edit(t, dir, bansJSON, broken)
|
edit(t, dir, bansJSON, broken)
|
||||||
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
|
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
|
||||||
wantTakenIn(t, lines, dir, clientsJSON)
|
wantTakenIn(t, lines, dir, clientsJSON)
|
||||||
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
|
||||||
// The next write sets it aside, logged with where the error is, and
|
// The next write sets it aside, logged with where the error is, and
|
||||||
// writes bans.json again from what smallwebwaf still holds.
|
// writes bans.json again from what smallwebwaf still holds.
|
||||||
@@ -861,7 +1136,14 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
|
|||||||
t.Errorf("set aside with %v", line)
|
t.Errorf("set aside with %v", line)
|
||||||
}
|
}
|
||||||
|
|
||||||
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
|
// It is raised as a file_error alert, with the same file and error.
|
||||||
|
waiting := params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||||
|
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
|
||||||
|
waiting[0].Detail["file"] != path+".bad" || waiting[0].Detail["error"] != message {
|
||||||
|
t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad")
|
||||||
|
}
|
||||||
|
|
||||||
|
wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
|
||||||
|
|
||||||
if got := readFile(t, path+".bad"); got != broken {
|
if got := readFile(t, path+".bad"); got != broken {
|
||||||
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
|
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
|
||||||
@@ -894,7 +1176,7 @@ func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) {
|
|||||||
wantTakenIn(t, lines, dir, bansJSON)
|
wantTakenIn(t, lines, dir, bansJSON)
|
||||||
|
|
||||||
wantMetric(t, scrape(t, params),
|
wantMetric(t, scrape(t, params),
|
||||||
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2)
|
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json",instance="app"}`, 2)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
|
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
|
||||||
@@ -959,7 +1241,7 @@ func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
wantMetric(t, scrape(t, params),
|
wantMetric(t, scrape(t, params),
|
||||||
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1)
|
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json",instance="app"}`, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
|
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
|
||||||
@@ -990,10 +1272,11 @@ func midnight() time.Time {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// newParams returns Params for the state files in dir, with parts that
|
// newParams returns Params for the state files in dir, with parts that
|
||||||
// hold nothing yet. GeoJS is never asked.
|
// hold nothing yet. GeoJS is never asked, and the alerts, at most two an
|
||||||
|
// hour, are never sent.
|
||||||
func newParams(dir string) state.Params {
|
func newParams(dir string) state.Params {
|
||||||
discard := slog.New(slog.DiscardHandler)
|
discard := slog.New(slog.DiscardHandler)
|
||||||
m := metrics.New(1)
|
m := metrics.New(1, "app")
|
||||||
|
|
||||||
return state.Params{
|
return state.Params{
|
||||||
Dir: dir,
|
Dir: dir,
|
||||||
@@ -1010,6 +1293,14 @@ func newParams(dir string) state.Params {
|
|||||||
GeoJS: lookup.New(lookup.Params{
|
GeoJS: lookup.New(lookup.Params{
|
||||||
Now: midnight, ProcessLog: discard, Metrics: m,
|
Now: midnight, ProcessLog: discard, Metrics: m,
|
||||||
}),
|
}),
|
||||||
|
Alerts: alerts.New(alerts.Params{
|
||||||
|
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||||
|
Events: alerts.Events(),
|
||||||
|
Cooldown: 15 * time.Minute,
|
||||||
|
MaxPerHour: 2,
|
||||||
|
Instance: "fsn1app1/gitea",
|
||||||
|
Now: midnight,
|
||||||
|
}),
|
||||||
Now: midnight,
|
Now: midnight,
|
||||||
ProcessLog: discard,
|
ProcessLog: discard,
|
||||||
Metrics: m,
|
Metrics: m,
|
||||||
@@ -1017,32 +1308,59 @@ func newParams(dir string) state.Params {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// fill puts a permanent ban an admin made, a ban for a broken limit and
|
// fill puts a permanent ban an admin made, a ban for a broken limit and
|
||||||
// one for a clear sign of attack, clients with counts and histories, and
|
// one for a clear sign of attack, clients with counts and histories,
|
||||||
// GeoJS answers into the parts of params.
|
// GeoJS answers, and alerts, as filledAlertsJSON holds them, into the
|
||||||
|
// parts of params.
|
||||||
func fill(params state.Params) {
|
func fill(params state.Params) {
|
||||||
now := midnight()
|
now := midnight()
|
||||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
params.Ledger.Load([]bans.Ban{permanentBan()})
|
params.Ledger.Load([]bans.Ban{permanentBan()})
|
||||||
params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1})
|
params.Ledger.BanForLimit(client, now, bans.Notes{
|
||||||
|
ASN: asn, ASName: asName, Country: "DE", Limit: 1,
|
||||||
|
})
|
||||||
params.Ledger.BanForAttack(netip.MustParsePrefix("192.0.2.1/32"), now,
|
params.Ledger.BanForAttack(netip.MustParsePrefix("192.0.2.1/32"), now,
|
||||||
bans.Notes{RuleID: "env-file", Target: "path"})
|
bans.Notes{RuleID: "env-file", Target: "path"})
|
||||||
|
|
||||||
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
|
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
|
||||||
params.Limiter.Count(netip.MustParsePrefix(c), now)
|
params.Limiter.Count(netip.MustParsePrefix(c), now, whole)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
params.Limiter.CountBytes(client, now, 8, whole)
|
||||||
params.Limiter.AddToHistory(client, now, ratelimit.Request{
|
params.Limiter.AddToHistory(client, now, ratelimit.Request{
|
||||||
Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
|
Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
|
||||||
})
|
})
|
||||||
|
params.Limiter.AddLookup(client, now.Add(-time.Hour), asn, asName, "DE")
|
||||||
|
|
||||||
params.GeoJS.Load([]lookup.Answer{
|
params.GeoJS.Load([]lookup.Answer{
|
||||||
{Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now},
|
{
|
||||||
|
Client: client, ASN: asn, ASName: asName, Country: "DE",
|
||||||
|
Answered: now.Add(-time.Hour), Used: now,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
Client: netip.MustParsePrefix("192.0.2.1/32"),
|
Client: netip.MustParsePrefix("192.0.2.1/32"),
|
||||||
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
|
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// An alert waiting, a repeat of it the cooldown holds back, another
|
||||||
|
// alert waiting, and one past the two an hour, for the hour's summary.
|
||||||
|
ban := alerts.Alert{
|
||||||
|
Event: alerts.EventBan, Client: client.Addr(), Netblock: client, Country: "DE",
|
||||||
|
Reason: "requests per minute over the limit of 1",
|
||||||
|
Detail: map[string]any{"cause": "limit"},
|
||||||
|
}
|
||||||
|
params.Alerts.Raise(ban)
|
||||||
|
params.Alerts.Raise(ban)
|
||||||
|
params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventFileError, Reason: "writing the state files failed",
|
||||||
|
Detail: map[string]any{
|
||||||
|
"file": "/var/lib/smallwebwaf/bans.json", "error": "no space left on device",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
params.Alerts.Raise(alerts.Alert{
|
||||||
|
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// permanentBan is the ban permanentBansJSON holds.
|
// permanentBan is the ban permanentBansJSON holds.
|
||||||
@@ -1053,6 +1371,8 @@ func permanentBan() bans.Ban {
|
|||||||
Cause: bans.CauseAdmin,
|
Cause: bans.CauseAdmin,
|
||||||
Reason: "scrapes every commit",
|
Reason: "scrapes every commit",
|
||||||
Notes: bans.Notes{
|
Notes: bans.Notes{
|
||||||
|
ASN: asn,
|
||||||
|
ASName: asName,
|
||||||
Country: "DE",
|
Country: "DE",
|
||||||
Limit: 1000,
|
Limit: 1000,
|
||||||
Window: "minute",
|
Window: "minute",
|
||||||
@@ -1098,13 +1418,13 @@ func wantLiftedBanKept(
|
|||||||
|
|
||||||
netblock := netip.MustParsePrefix(liftedClient + "/32")
|
netblock := netip.MustParsePrefix(liftedClient + "/32")
|
||||||
|
|
||||||
_, banned := ledger.Check(netblock.Addr(), afterLifting())
|
_, banned, _ := ledger.Check(netblock.Addr(), afterLifting())
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the lifted ban refuses")
|
t.Error("the lifted ban refuses")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Were the lifted ban counted, the next would last three hours.
|
// Were the lifted ban counted, the next would last three hours.
|
||||||
ban := ledger.BanForLimit(netblock, afterLifting(), bans.Notes{})
|
ban, _ := ledger.BanForLimit(netblock, afterLifting(), bans.Notes{})
|
||||||
if ban.Expires.Sub(ban.Start) != time.Hour {
|
if ban.Expires.Sub(ban.Start) != time.Hour {
|
||||||
t.Errorf("the next ban lasts %s, want 1h", ban.Expires.Sub(ban.Start))
|
t.Errorf("the next ban lasts %s, want 1h", ban.Expires.Sub(ban.Start))
|
||||||
}
|
}
|
||||||
@@ -1350,8 +1670,8 @@ func scrape(t *testing.T, params state.Params) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// metric returns the value of series in text, the metrics, such as
|
// metric returns the value of series in text, the metrics, such as
|
||||||
// smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test
|
// smallwebwaf_state_file_writes_total{file="bans.json",instance="app"}, or
|
||||||
// if there is no such series.
|
// fails the test if there is no such series.
|
||||||
func metric(t *testing.T, text, series string) float64 {
|
func metric(t *testing.T, text, series string) float64 {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
@@ -1380,7 +1700,7 @@ func wantWriteFailed(t *testing.T, params state.Params, name string) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
got := scrape(t, params)
|
got := scrape(t, params)
|
||||||
file := `{file="` + name + `"}`
|
file := `{file="` + name + `",instance="app"}`
|
||||||
|
|
||||||
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1)
|
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1)
|
||||||
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1)
|
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1)
|
||||||
|
|||||||
+7
-5
@@ -7,9 +7,10 @@
|
|||||||
# request bans for good, that `sv stop` stops smallwebwaf in order, that
|
# request bans for good, that `sv stop` stops smallwebwaf in order, that
|
||||||
# `docker stop` stops the container without having to kill it, and that
|
# `docker stop` stops the container without having to kill it, and that
|
||||||
# a new container on the same volume still refuses the banned client. The
|
# a new container on the same volume still refuses the banned client. The
|
||||||
# containers, the volume and both images are removed however the script
|
# containers run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent
|
||||||
# ends. Building the app needs network access, for nixpkgs' binary cache.
|
# to GeoJS. The containers, the volume and both images are removed however
|
||||||
# script/check does not run this.
|
# the script ends. Building the app needs network access, for nixpkgs'
|
||||||
|
# binary cache. script/check does not run this.
|
||||||
set -eu
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -63,12 +64,13 @@ logged() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
# start_container: run the app's container, with the state files on the
|
# start_container: run the app's container, with the state files on the
|
||||||
# volume and a rate limit of one request a minute, and wait until it is
|
# volume, a rate limit of one request a minute and no client looked up,
|
||||||
# healthy.
|
# and wait until it is healthy.
|
||||||
start_container() {
|
start_container() {
|
||||||
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
||||||
--volume "$VOLUME:/var/lib/smallwebwaf" \
|
--volume "$VOLUME:/var/lib/smallwebwaf" \
|
||||||
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
|
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
|
||||||
|
--env SWWAF_LOOKUP_SOURCE=off \
|
||||||
"$APP_IMAGE" >/dev/null
|
"$APP_IMAGE" >/dev/null
|
||||||
wait_for "the health check did not pass" healthy
|
wait_for "the health check did not pass" healthy
|
||||||
address="$(docker port "$CONTAINER" 8080/tcp)"
|
address="$(docker port "$CONTAINER" 8080/tcp)"
|
||||||
|
|||||||
Executable
+19
@@ -0,0 +1,19 @@
|
|||||||
|
#!/bin/sh
|
||||||
|
# script/tidy: write go.mod and go.sum as `go mod tidy` writes them, which
|
||||||
|
# the test phase of the Dockerfile checks. This builds the Dockerfile's
|
||||||
|
# tidy-files stage, which holds the two files alone, and --output writes
|
||||||
|
# them into the working tree. The build makes no image, so it has no tag.
|
||||||
|
# --no-cache because `go mod tidy` asks the module proxy, whose answers a
|
||||||
|
# cached layer would repeat.
|
||||||
|
set -eu
|
||||||
|
|
||||||
|
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
|
||||||
|
|
||||||
|
main() {
|
||||||
|
cd "$ROOT"
|
||||||
|
docker build --no-cache \
|
||||||
|
--target tidy-files \
|
||||||
|
--output "type=local,dest=$ROOT" .
|
||||||
|
}
|
||||||
|
|
||||||
|
main "$@"
|
||||||
Reference in New Issue
Block a user