Compare commits

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

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

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

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

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

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

Model: opus-5-5
2026-10-07 01:13:16 +02:00
clawbot bff65f4e2f Settings given as files: the _FILE form of every setting (closes #87)
check / check (push) Waiting to run
Every setting X may instead be given as a file that X_FILE names, read
once at start: its contents, less one trailing newline, are the value,
checked as X would be. X and X_FILE both set, or a file that cannot be
read, stops the start with a message naming the variable. The logged
settings name the file, and mask a token read from one.
SWWAF_LOG_REMOTE_TLS_CA_FILE, whose value is a file already, has no
_FILE form. The health check reads only SWWAF_LISTEN_ADDR and
SWWAF_UPSTREAM_URL, so no other setting or file can fail it.

Judgement call: an invalid value read from a file is named as X, not X_FILE.
Rule suppressed: gosec G304 on reading the named file, as for the CA file.

Model: opus-5-5
2026-10-07 00:01:19 +02:00
clawbot ee9ba08a8a Bans an admin makes or lifts: the admin cause, a reason, lifted bans kept (closes #86)
check / check (push) Waiting to run
A bans.json entry without a cause gets the cause admin, written back so.
Bans whose cause is admin are never dropped and do not count toward
SWWAF_MAX_BANS, so setting a ban's cause to admin keeps it. Bans
smallwebwaf makes get a reason: the limit broken or the rule matched. A
lifted ban refuses nothing, is kept, and makes no later ban longer.
smallwebwaf_bans_made_total counts admin bans an edit adds while running;
earlier_bans counts admin in place of without_cause.

Judgement call: lifted lifts at once, whatever time it gives.
Judgement call: a lifted ban still counts in earlier_bans.
Known gap: a ban dropped from behind an admin's ban on its netblock leaves that netblock's later earlier_bans.

Model: opus-5-5
2026-10-06 23:09:52 +02:00
clawbot 0797e5def2 Send every log line to a syslog server as well (closes #28)
check / check (push) Waiting to run
With SWWAF_LOG_REMOTE_URL set (syslog+udp, syslog+tcp or syslog+tls),
every line on stdout is also sent as the message of an RFC 5424 record,
octet-counted over TCP and TLS, from a bounded buffer that drops its
oldest line when full, so a slow or unreachable server holds up nothing.
Failed connections are retried with backoff; lines sent, dropped and
waiting are metrics. At a stop the lines still waiting get at most two
seconds. SWWAF_LOG_REMOTE_APP_NAME defaults to SWWAF_INSTANCE_NAME; while
sending, an app name RFC 5424 does not allow stops the start. Standard
library only: log/syslog writes only the older format.

Model: opus-5-5
2026-10-06 22:02:36 +02:00
clawbot e77dfb6891 Rule files, and bans for a clear sign of attack (closes #24)
check / check (push) Waiting to run
Every *.rules file in SWWAF_RULES_DIR not named with a leading dot is
read at start, and again 2 seconds after the directory's last change.
Each request is checked against the rules after the rate limits: log
notes a match, block refuses with 403, ban refuses and bans the netblock
for SWWAF_ATTACK_BAN_DURATION, made permanent by its next request or
attack. path, query and uri are matched as the request line sent them;
header:Host and header:Transfer-Encoding are refused. Bans gain a cause.
The image ships 00-default.rules.

Judgement call: a header sent twice is matched with its values joined
by ", ".
Judgement call: SWWAF_MAX_BAN_DURATION does not cap a ban for an attack.
Not in this unit: offences for rule matches, with the error burst.

Model: opus-5-5
2026-10-06 20:38:35 +02:00
clawbot 74bdc6a449 Leave SWWAF_RATE_LIMIT_EXEMPT_PATHS out of the request rate limits (closes #77)
check / check (push) Successful in 3m2s
A request is neither counted nor refused by the request rate limits
when its path as sent, the path the app receives, not percent-decoded,
starts with one of the comma-separated prefixes in
SWWAF_RATE_LIMIT_EXEMPT_PATHS, so /%61ssets/x is not under /assets/. A
request whose decoded path contains .. or a backslash, or whose path as
sent holds an encoded slash, is never exempt, since an app may act on
it as a path outside every prefix, such as /assets/..%2Flogin as
/login. The static lists, bans and the country lists still apply, and
its log line has no counts. The setting is empty by default, and a
prefix that does not start with / stops the start. README.md documents
it.

Model: opus-5-5
2026-10-06 18:47:13 +02:00
clawbot 808e69f442 Log the rest of the request log's fields (closes #79)
check / check (push) Successful in 3m23s
Each request log line now has the fields "Request log" in SPEC.md lists
whose features are built: instance (SWWAF_INSTANCE_NAME), scheme,
request_id (a trusted proxy's X-Request-ID or a new one, sent on to the
app), forwarded_for, client_group, content_type, content_length, the
headers SWWAF_LOG_REQUEST_HEADERS names, has_authorization, has_cookie,
websocket, response_content_type, cache_control, location, counts and
the timings. Authorization, Cookie and Set-Cookie values are never
logged. An entry of SWWAF_LOG_REQUEST_HEADERS that is not a header name,
or is Host or Transfer-Encoding, stops the start.

Deviation: counts has request totals only.
Deviation: SWWAF_INSTANCE_NAME is on request lines only.

Model: opus-5-5
2026-10-06 17:26:21 +02:00
clawbot 6ec52e5b87 Take in an admin's edits of the state files while running (closes #68)
check / check (push) Successful in 4m1s
smallwebwaf watches SWWAF_STATE_DIR with fsnotify and takes in a saved
edit of a state file in place of what it held. It knows its own writes
by the SHA-256 of what it last read or wrote; each write first takes in
an edit made since. An edit that does not parse is renamed to
<name>.bad at the next write. Each edit taken in or set aside is logged
and counted. Every ban on a netblock is checked, and the next ban is
worked out from the one that ended last. README.md says how to add and
lift a ban.

Judgement call: a broken edit is set aside at the next write, since an
editor's file can be read half written.

Model: opus-5-5
2026-10-06 14:18:13 +02:00
clawbot cff385af41 Observe mode: log what would be refused, refuse nothing (closes #78)
check / check (push) Successful in 3m54s
SWWAF_MODE (default enforce) takes enforce or observe. In observe mode a
request that SWWAF_DENY_NETS, a ban, the country lists or a rate limit
would refuse is passed to the app, and its log line names that refusal
in would_action. The size and time limits and the 401 still apply. A
broken limit makes no ban; bans read from bans.json are kept but refuse
nothing, and Ledger.Find reads them without counting a refusal in their
notes.

Judgement call: in observe mode a broken limit does not reset the
client's counters, since the reset comes with the ban.
Judgement call: a request a ban would refuse keeps ban_expires.

Model: opus-5-5
2026-10-06 13:04:43 +02:00
clawbot 234c5eac60 Serve Prometheus metrics behind SWWAF_METRICS_TOKEN (closes #23)
check / check (push) Successful in 3m21s
GET /_smallwebwaf/metrics answers in the Prometheus text format for a
request carrying SWWAF_METRICS_TOKEN, 401 without it and 404 while it is
unset. Every request under /_smallwebwaf/ but the health check now goes
through the checks and is answered where it would be forwarded, 404 for
any path but the metrics, so none reaches the app. In the client's
history a 401 counts as refused, the metrics and the 404s as neither.
SWWAF_METRICS_TOP_N bounds the series by country, the rest counted as
other.

Deviation: go.mod and go.sum written by hand, as go runs only through
make.
Deviation: no metrics yet for state files read again after an edit or
edits set aside; that work is not merged.

Model: opus-5-5
2026-10-06 11:40:27 +02:00
clawbot 68f687cb0c Run the GeoJS lookup tests on a clock the test controls (closes #73)
check / check (push) Successful in 3m32s
The tests that have GeoJS asked now run in a synctest bubble, so a wait
lasts exactly as long as it should however slowly the test process runs:
a new client's wait is checked to be exactly one second, and its next
request exactly no wait. A request waiting on the network would stop the
bubble's clock, so the stand-in for GeoJS now answers in place of the
network, through a transport that a test-only file lets the tests set.
The test of the failure log uses an abandoned request instead of a
closed port.

Model: opus-5-5
2026-10-06 09:12:50 +02:00
clawbot df2c5042d2 Keep the bans, the clients and GeoJS's answers in state files (closes #17)
check / check (push) Successful in 3m24s
smallwebwaf now copies its state to bans.json, clients.json and
lookups.json in SWWAF_STATE_DIR, as "Persistent state" in SPEC.md
describes, and reads them back at start, so a restart lifts no ban and
gives no client a fresh allowance. Each client gains a history, and a
ban's notes count the netblock's requests. bans.json is written
SWWAF_STATE_WRITE_DELAY after a ban, and every file every
SWWAF_STATE_COUNTER_INTERVAL and at the stop. A ban read back is masked
to its netblock and refuses every client in it. A file that does not
parse, an unknown version, an entry without a field it needs, or an
unwritable directory stops the start.

Deviation: no AS number or name, and no ban cause, reason or lifting yet.

Model: opus-5-5
2026-10-06 08:31:52 +02:00
clawbot 73ca94f850 Ban the netblock of a client that breaks a rate limit, in memory (closes #18)
check / check (push) Successful in 3m48s
A request over a rate limit is refused with SWWAF_BAN_RESPONSE and bans
the client's netblock: an hour at first, three times the last ban when
broken again within a day of its end, permanent past seven days. The
ban ledger in internal/bans is checked after the static lists and
before the lookup, and the requests it refuses are not counted. A ban
resets the client's counters and carries notes holding the request
that broke the limit, as SPEC.md now says. At most SWWAF_MAX_BANS are
held. SWWAF_BAN_RESPONSE also answers SWWAF_DENY_NETS and the country
lists.

Judgement call: the six ban settings cannot be off.
Judgement call: a permanent ban's ban_expires is "permanent".

Model: opus-5-5
2026-10-06 05:29:03 +02:00
clawbot 0f85c9ae07 The header size and the idle time as settings (closes #70)
check / check (push) Successful in 4m56s
SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES (default 32K) and
SWWAF_CLIENT_IDLE_TIMEOUT (default 120s) replace the two values the
proxy fixed. The idle time is read like the other durations, and can
be off.

Go's server reads 4K past the header limit it is given before it
refuses, so it is still given the setting less 4K. The header size
must be more than 4K and cannot be off; any other value stops the
start with a message that does not offer off.

SPEC.md and README.md say so. README.md lists both settings, no longer
calls them fixed, and names them as built.

Model: opus-5-5
2026-10-06 05:05:31 +02:00
66 changed files with 17896 additions and 969 deletions
+9
View File
@@ -162,6 +162,15 @@ RUN groupadd --system --gid 65532 smallwebwaf \
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \ --gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
smallwebwaf smallwebwaf
# The state files' directory, SWWAF_STATE_DIR by default, where a volume
# is mounted to keep them across deploys. The run script gives it to the
# smallwebwaf user at each start.
RUN mkdir /var/lib/smallwebwaf
# The default rule file, in SWWAF_RULES_DIR by default, where an app's
# Dockerfile can copy rule files of its own beside it.
COPY share/rules.d/00-default.rules /etc/smallwebwaf/rules.d/00-default.rules
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv # runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
# looks too. # looks too.
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
+930 -123
View File
File diff suppressed because it is too large Load Diff
+8 -4
View File
@@ -32,7 +32,7 @@ from a directory of hand-editable text files.
- Defence against traffic floods that saturate the host's network link. That - Defence against traffic floods that saturate the host's network link. That
needs help upstream of the host. needs help upstream of the host.
- A web UI or a configuration file. Settings are environment variables. Apart - A web UI or a configuration file. Settings are environment variables. Apart
from settings given as files (the `_FILE` form of any setting, such as from settings given as files (the `_FILE` form of a setting, such as
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state `SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
files and the lookup database, the only files read are the rule files, which files and the lookup database, the only files read are the rule files, which
hold one regex per line and nothing more elaborate. hold one regex per line and nothing more elaborate.
@@ -293,11 +293,13 @@ it.
needs: an alert destination, an account key, a token. needs: an alert destination, an account key, a token.
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its - Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
container, and so its environment variables, with the app it protects. container, and so its environment variables, with the app it protects.
- Any limit or threshold can be switched off with the value `off`. - Any limit or threshold can be switched off with the value `off`, except
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`.
- A list set to an empty value is an empty list, and replaces the default. - A list set to an empty value is an empty list, and replaces the default.
- Every setting may instead be given as a file holding the value, named by the - Every setting may instead be given as a file holding the value, named by the
setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for
secrets and long lists. secrets and long lists. `SWWAF_LOG_REMOTE_TLS_CA_FILE`, whose value names a
file already, has no `_FILE` form.
- Settings, including those given as files, are read once at start; changing one - Settings, including those given as files, are read once at start; changing one
means restarting the container. The files `smallwebwaf` watches while it runs means restarting the container. The files `smallwebwaf` watches while it runs
are its state files, its rule files and the lookup database. are its state files, its rule files and the lookup database.
@@ -413,7 +415,9 @@ The settings, by group:
headers, its body. headers, its body.
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest - `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
request line and headers a client may send. Over it, `smallwebwaf` answers request line and headers a client may send. Over it, `smallwebwaf` answers
`431` and closes the connection, and nothing reaches the app. `431` and closes the connection, and nothing reaches the app. It must be
more than `4K`, and cannot be `off`: Go's HTTP server always has such a
limit, and reads 4 KiB past the one it is given before it refuses.
- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open - `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open
connection may wait for its next request before `smallwebwaf` closes it. connection may wait for its next request before `smallwebwaf` closes it.
It is longer than the 90 seconds after which traefik, by default, closes a It is longer than the 90 seconds after which traefik, by default, closes a
+17 -1
View File
@@ -2,4 +2,20 @@ module sneak.berlin/go/smallwebwaf
go 1.26.0 go 1.26.0
require github.com/hashicorp/golang-lru/v2 v2.0.7 require (
github.com/fsnotify/fsnotify v1.10.1
github.com/hashicorp/golang-lru/v2 v2.0.7
github.com/prometheus/client_golang v1.24.1
)
require (
github.com/beorn7/perks v1.0.1 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/kylelemons/godebug v1.1.0 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/prometheus/client_model v0.6.2 // indirect
github.com/prometheus/common v0.70.1 // indirect
github.com/prometheus/procfs v0.21.1 // indirect
golang.org/x/sys v0.47.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect
)
+38
View File
@@ -1,2 +1,40 @@
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/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
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/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/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
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/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
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/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
github.com/prometheus/client_golang v1.24.1/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/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY=
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
go.uber.org/goleak v1.3.0/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/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
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=
+882
View File
@@ -0,0 +1,882 @@
// Package alerts sends alerts on bans, on a source that fails and on a
// file with an error to each destination set: to the webhook
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
// webhook schema" section of SPEC.md describes, to the Slack incoming
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
// ntfy topic SWWAF_ALERT_NTFY_URL names. A repeat within
// SWWAF_ALERT_COOLDOWN is held back, and so is an alert past
// SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary. The others wait in a
// bounded queue of each destination's own, so that a destination that is
// slow or unreachable holds up neither the others nor any request. The
// state is written to alerts.json and read from it by the state package.
// Nothing logged names a destination's URL, whose path or query can carry
// a secret.
package alerts
import (
"bytes"
"cmp"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net/http"
"net/netip"
"net/url"
"slices"
"strings"
"sync"
"sync/atomic"
"time"
)
// The events an alert is for, as SWWAF_ALERT_EVENTS names them.
const (
// EventBan is a ban smallwebwaf made.
EventBan = "ban"
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
// made permanent.
EventPermanentBan = "permanent_ban"
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
// Core Rule Set, the anomaly thresholds and the reputation sources;
// nothing raises them yet.
EventWAFBlock = "waf_block"
EventAnomaly = "anomaly"
EventReputationHit = "reputation_hit"
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
EventSourceFailure = "source_failure"
// EventFileError is a rule file or state file edited while smallwebwaf
// runs that does not parse, or a state file that cannot be written.
EventFileError = "file_error"
// EventSummary is the summary of the alerts an hour held back past
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
EventSummary = "summary"
)
// Events returns every event SWWAF_ALERT_EVENTS can name, which is its
// default.
func Events() []string {
return []string{
EventBan, EventPermanentBan, EventWAFBlock, EventAnomaly,
EventReputationHit, EventSourceFailure, EventFileError,
}
}
// The destinations alerts are sent to, as the metrics and alerts.json
// name them.
const (
// DestinationWebhook is the webhook SWWAF_ALERT_WEBHOOK_URL names.
DestinationWebhook = "webhook"
// DestinationSlack is the Slack incoming webhook
// SWWAF_ALERT_SLACK_WEBHOOK_URL names.
DestinationSlack = "slack"
// DestinationNtfy is the ntfy topic SWWAF_ALERT_NTFY_URL names.
DestinationNtfy = "ntfy"
)
// Destinations returns every destination alerts can be sent to.
func Destinations() []string {
return []string{DestinationWebhook, DestinationSlack, DestinationNtfy}
}
const (
// queueSize is the most alerts that wait to be sent to a destination.
// Past it, the oldest is dropped.
queueSize = 1000
// sendTimeout bounds one request to a destination.
sendTimeout = 10 * time.Second
// After a request to a destination fails, the alert is sent again a
// second later, and retryDelayFactor times as long after each further
// failure in a row, up to a minute.
firstRetryDelay = time.Second
retryDelayFactor = 2
maxRetryDelay = time.Minute
// maxAnswerBytes is the most of a destination's answer that is read.
maxAnswerBytes = 64 << 10
)
var (
errStatus = errors.New("the destination answered")
// errRefused is a 4xx answer other than 408 and 429: the destination
// refuses the alert itself, and would refuse it again.
errRefused = errors.New("the destination refused the alert, answering")
)
// Params are what New needs. With none of WebhookURL, SlackURL and
// NtfyURL set, no alert is sent.
type Params struct {
// WebhookURL is where each alert is posted as JSON
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. WebhookHeaders
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
WebhookURL *url.URL
WebhookHeaders http.Header
// SlackURL is the Slack incoming webhook each alert is posted to as a
// message (SWWAF_ALERT_SLACK_WEBHOOK_URL), nil while it is unset.
SlackURL *url.URL
// NtfyURL is the ntfy topic each alert is published to
// (SWWAF_ALERT_NTFY_URL), nil while it is unset. NtfyToken, unless
// empty, is sent with each as a bearer token (SWWAF_ALERT_NTFY_TOKEN).
NtfyURL *url.URL
NtfyToken string
// Events are the events alerts are sent for (SWWAF_ALERT_EVENTS).
Events []string
// Cooldown is how long a repeat of an alert is held back
// (SWWAF_ALERT_COOLDOWN), 0 for no time. MaxPerHour is the most alerts
// sent in an hour (SWWAF_ALERT_MAX_PER_HOUR), 0 for no limit.
Cooldown time.Duration
MaxPerHour int
// Instance is SWWAF_INSTANCE_NAME, which every alert gives.
Instance string
// Now tells the time of an alert, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives the requests to a destination that fail.
ProcessLog *slog.Logger
}
// Alert is one alert, as the webhook is sent it and alerts.json holds it,
// with the fields of the "Alert webhook schema" section of SPEC.md. ASN
// and ASName are empty until AS numbers are looked up.
//
//nolint:tagliatelle // SPEC.md's alert webhook schema names its fields in snake_case
type Alert struct {
Instance string `json:"instance"`
Time time.Time `json:"time"`
Event string `json:"event"`
Client netip.Addr `json:"client"`
Netblock netip.Prefix `json:"netblock"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"`
// Reason is a short sentence, and Detail what is particular to the
// event: for a file_error, its "file", and for a source_failure, its
// "source", which the cooldown tells repeats by.
Reason string `json:"reason"`
Detail map[string]any `json:"detail"`
// SuppressedRepeats is how many repeats of the alert the cooldown
// held back since the last one let through.
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Cooldown is, for an event on a netblock, or about a file or a source,
// when the last alert let through was raised, and how many repeats the
// cooldown has held back since, as alerts.json holds it.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Cooldown struct {
Event string `json:"event"`
Netblock netip.Prefix `json:"netblock"`
File string `json:"file,omitempty"`
Source string `json:"source,omitempty"`
Sent time.Time `json:"sent"`
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Hour is the hour under way, by the clock, as alerts.json holds it: when
// it started, how many alerts were let through in it, and how many were
// held back in it past MaxPerHour, by event, for its summary.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Hour struct {
Start time.Time `json:"start"`
Sent int `json:"sent"`
HeldBack map[string]int `json:"held_back"`
}
// State is what alerts.json holds: the cooldowns, the hour under way, and
// for each destination set, the alerts waiting to be sent to it, oldest
// first.
type State struct {
Cooldowns []Cooldown `json:"cooldowns"`
Hour Hour `json:"hour"`
Waiting map[string][]Alert `json:"waiting"`
}
// Counts are, for a destination, how many alerts it took, how many
// requests to it failed, and how many alerts were dropped from its full
// queue or given up as it refused them.
type Counts struct {
Sent, Failed, Dropped int64
}
// Queue takes the alerts raised, holds back those it must, and sends the
// others to each destination set, from a queue of the destination's own.
// It is safe for concurrent use.
type Queue struct {
params Params
// destinations are the destinations set, in the order of
// Destinations.
destinations []*destination
mu sync.Mutex
// cooldowns are the alerts last let through, by event and netblock,
// file or source.
cooldowns map[cooldownKey]*Cooldown
hour Hour
suppressed atomic.Int64
}
// destination is a destination set, with the alerts waiting to be sent
// to it. Its mu is taken after the Queue's, never before.
type destination struct {
// name is how the metrics and alerts.json name the destination, and
// setting the setting that is its URL, which the log names in place
// of the URL.
name string
setting string
url *url.URL
// message returns the body an alert is posted with, and the headers
// sent with it.
message func(alert *Alert) ([]byte, http.Header, error)
// httpClient follows no redirect: a redirect is a failure.
httpClient *http.Client
processLog *slog.Logger
// queued receives a value when an alert joins the queue, unless one
// waits already, so that run looks at the queue again.
queued chan struct{}
mu sync.Mutex
// waiting are the alerts waiting to be sent, oldest first.
waiting []*Alert
sent, failed, dropped atomic.Int64
}
// cooldownKey is what makes an alert a repeat of another: the same event
// on the same netblock, and about the same file or source, as its detail
// names them. Each is empty for an alert without one.
type cooldownKey struct {
event string
netblock netip.Prefix
file string
source string
}
// cooldownKeyOf returns what makes another alert a repeat of alert.
func cooldownKeyOf(alert *Alert) cooldownKey {
file, _ := alert.Detail["file"].(string)
source, _ := alert.Detail["source"].(string)
return cooldownKey{alert.Event, alert.Netblock, file, source}
}
// New returns a Queue with no alert yet.
func New(params Params) *Queue {
q := &Queue{
params: params,
cooldowns: map[cooldownKey]*Cooldown{},
hour: Hour{HeldBack: map[string]int{}},
}
if params.WebhookURL != nil {
q.addDestination(DestinationWebhook, "SWWAF_ALERT_WEBHOOK_URL",
params.WebhookURL, q.webhookMessage)
}
if params.SlackURL != nil {
q.addDestination(DestinationSlack, "SWWAF_ALERT_SLACK_WEBHOOK_URL",
params.SlackURL, slackMessage)
}
if params.NtfyURL != nil {
q.addDestination(DestinationNtfy, "SWWAF_ALERT_NTFY_URL",
params.NtfyURL, q.ntfyMessage)
}
return q
}
// Raise sends alert, which names its event and what is particular to it,
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event
// out. It gives alert the instance and the time. An alert that repeats
// the last one let through less than Cooldown before is held back and
// counted, and the next one let through gives that count. Past MaxPerHour
// alerts let through in the hour under way, by the clock, an alert is
// held back for that hour's summary instead, which is sent once the hour
// has ended; it starts no cooldown, and the repeats held back before it
// are given by the next alert let through. Raise never waits: an alert
// let through joins the queue of each destination, from which Run sends
// it, and with queueSize alerts waiting for a destination, the oldest is
// dropped.
func (q *Queue) Raise(alert Alert) {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
return
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
alert.Instance = q.params.Instance
alert.Time = now
if q.repeat(&alert, now) {
q.suppressed.Add(1)
return
}
q.endHour(now)
if q.params.MaxPerHour > 0 && q.hour.Sent >= q.params.MaxPerHour {
q.hour.HeldBack[alert.Event]++
q.suppressed.Add(1)
return
}
q.startCooldown(&alert, now)
q.hour.Sent++
q.queue(&alert)
}
// WouldSend reports whether Raise would let an alert for event on
// netblock through now: a destination is set, SWWAF_ALERT_EVENTS chooses
// event, no alert for event on netblock was let through less than
// Cooldown before, and fewer than MaxPerHour alerts have been let through
// in the hour under way. Unlike Raise, it counts nothing.
func (q *Queue) WouldSend(event string, netblock netip.Prefix) bool {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, event) {
return false
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
last, found := q.cooldowns[cooldownKey{event: event, netblock: netblock}]
if q.params.Cooldown > 0 && found && now.Sub(last.Sent) < q.params.Cooldown {
return false
}
q.endHour(now)
return q.params.MaxPerHour == 0 || q.hour.Sent < q.params.MaxPerHour
}
// Run sends the alerts waiting to each destination, from its own queue,
// as destination.run does, until ctx is done. It also ends each hour as
// Raise does, so that the hour's summary is sent as it ends. With no
// destination set, it returns at once.
func (q *Queue) Run(ctx context.Context) {
if len(q.destinations) == 0 {
return
}
var sending sync.WaitGroup
for _, d := range q.destinations {
sending.Go(func() { d.run(ctx) })
}
for {
q.mu.Lock()
untilHourEnds := q.hour.Start.Add(time.Hour).Sub(q.params.Now())
q.mu.Unlock()
select {
case <-ctx.Done():
sending.Wait()
return
case <-time.After(untilHourEnds):
q.mu.Lock()
q.endHour(q.params.Now())
q.mu.Unlock()
}
}
}
// Counts returns the counts of the destination name, all 0 for one not
// set.
func (q *Queue) Counts(name string) Counts {
for _, d := range q.destinations {
if d.name == name {
return Counts{
Sent: d.sent.Load(), Failed: d.failed.Load(), Dropped: d.dropped.Load(),
}
}
}
return Counts{}
}
// DestinationsSet returns the destinations set, in the order of
// Destinations.
func (q *Queue) DestinationsSet() []string {
names := make([]string, 0, len(q.destinations))
for _, d := range q.destinations {
names = append(names, d.name)
}
return names
}
// Suppressed is how many alerts were held back: by the cooldown, and past
// MaxPerHour. No destination is sent such an alert.
func (q *Queue) Suppressed() int64 {
return q.suppressed.Load()
}
// Snapshot returns the queue's state, as alerts.json holds it, with the
// cooldowns sorted by netblock, then by event, file and source.
func (q *Queue) Snapshot() State {
q.mu.Lock()
defer q.mu.Unlock()
state := State{
Cooldowns: make([]Cooldown, 0, len(q.cooldowns)),
Hour: q.hour,
Waiting: map[string][]Alert{},
}
state.Hour.HeldBack = maps.Clone(q.hour.HeldBack)
for _, cooldown := range q.cooldowns {
state.Cooldowns = append(state.Cooldowns, *cooldown)
}
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
})
for _, d := range q.destinations {
state.Waiting[d.name] = d.snapshot()
}
return state
}
// Load puts state, read from alerts.json, in place of the queue's state.
// Each cooldown's netblock is masked to its length, so that
// 203.0.113.9/24 is 203.0.113.0/24. The alerts waiting for a destination
// that is not set are dropped, and so are the oldest past queueSize
// alerts waiting for one that is.
func (q *Queue) Load(state State) {
q.mu.Lock()
defer q.mu.Unlock()
q.cooldowns = map[cooldownKey]*Cooldown{}
for _, cooldown := range state.Cooldowns {
cooldown.Netblock = cooldown.Netblock.Masked()
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
q.cooldowns[key] = &cooldown
}
q.hour = state.Hour
q.hour.HeldBack = maps.Clone(state.Hour.HeldBack)
if q.hour.HeldBack == nil {
q.hour.HeldBack = map[string]int{}
}
for _, d := range q.destinations {
d.load(state.Waiting[d.name])
}
}
// addDestination adds a destination: its name, the setting that gives
// its URL, that URL, target, and message, which makes the messages sent
// to it.
func (q *Queue) addDestination(
name, setting string, target *url.URL,
message func(alert *Alert) ([]byte, http.Header, error),
) {
q.destinations = append(q.destinations, &destination{
name: name,
setting: setting,
url: target,
message: message,
httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
},
processLog: q.params.ProcessLog,
queued: make(chan struct{}, 1),
})
}
// repeat reports whether alert, raised at now, repeats the last one let
// through less than Cooldown before, and counts it if it does.
func (q *Queue) repeat(alert *Alert, now time.Time) bool {
if q.params.Cooldown == 0 {
return false
}
last, found := q.cooldowns[cooldownKeyOf(alert)]
if !found || now.Sub(last.Sent) >= q.params.Cooldown {
return false
}
last.SuppressedRepeats++
return true
}
// startCooldown gives alert, let through at now, the count of the repeats
// held back since the last one let through, and notes alert as the last
// one let through.
func (q *Queue) startCooldown(alert *Alert, now time.Time) {
if q.params.Cooldown == 0 {
return
}
key := cooldownKeyOf(alert)
last, found := q.cooldowns[key]
if found {
alert.SuppressedRepeats = last.SuppressedRepeats
}
q.cooldowns[key] = &Cooldown{
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
Sent: now,
}
}
// endHour ends the hour under way, if now is past it: it queues that
// hour's summary when alerts were held back in it past MaxPerHour, and
// forgets the cooldowns that have run out with no repeat held back, which
// no alert needs any more.
func (q *Queue) endHour(now time.Time) {
start := now.Truncate(time.Hour)
if !start.After(q.hour.Start) {
return
}
heldBack := 0
for _, count := range q.hour.HeldBack {
heldBack += count
}
if heldBack > 0 {
q.queue(&Alert{
Instance: q.params.Instance,
Time: now,
Event: EventSummary,
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
Detail: map[string]any{
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
},
})
}
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
for key, cooldown := range q.cooldowns {
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
delete(q.cooldowns, key)
}
}
}
// queue adds alert to the alerts waiting for each destination.
func (q *Queue) queue(alert *Alert) {
for _, d := range q.destinations {
d.add(alert)
}
}
// run sends the alerts waiting, oldest first, until ctx is done. An alert
// stays in the queue until the destination answers it with a 2xx status,
// or refuses it with a 4xx status other than 408 and 429: a refused alert
// is logged, counted as dropped, and given up, so that the next is sent.
// Any other request that fails is logged, and the alert sent again
// firstRetryDelay later, retryDelayFactor times as long after each
// further failure in a row, up to maxRetryDelay.
func (d *destination) run(ctx context.Context) {
var (
retryDelay time.Duration
retryAt time.Time
)
for {
alert := d.oldest()
var due <-chan time.Time // nil while no alert waits
if alert != nil {
due = time.After(time.Until(retryAt))
}
select {
case <-ctx.Done():
return
case <-d.queued:
case <-due:
err := d.send(ctx, alert)
switch {
case err == nil:
d.remove(alert)
d.sent.Add(1)
retryDelay = 0
retryAt = time.Time{}
case errors.Is(err, errRefused):
d.remove(alert)
d.failed.Add(1)
d.dropped.Add(1)
retryDelay = 0
retryAt = time.Time{}
d.processLog.Warn("gave up an alert "+d.setting+" refused",
"event", alert.Event, "error", err.Error())
case ctx.Err() == nil: // not cut off as smallwebwaf stops
d.failed.Add(1)
retryDelay = min(max(retryDelayFactor*retryDelay, firstRetryDelay),
maxRetryDelay)
retryAt = time.Now().Add(retryDelay)
d.processLog.Warn("sending an alert to "+d.setting+" failed",
"error", err.Error(), "sending_again_in", retryDelay.String())
}
}
}
}
// add adds alert to the alerts waiting, first dropping the oldest while
// queueSize wait, and has run look at the queue again.
func (d *destination) add(alert *Alert) {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) == queueSize {
d.waiting = slices.Delete(d.waiting, 0, 1)
d.dropped.Add(1)
}
d.waiting = append(d.waiting, alert)
select {
case d.queued <- struct{}{}:
default: // a value waits already
}
}
// load puts waiting, read from alerts.json, in place of the alerts
// waiting, as add adds them.
func (d *destination) load(waiting []Alert) {
d.mu.Lock()
d.waiting = nil
d.mu.Unlock()
for _, alert := range waiting {
d.add(&alert)
}
}
// snapshot returns the alerts waiting, oldest first.
func (d *destination) snapshot() []Alert {
d.mu.Lock()
defer d.mu.Unlock()
waiting := make([]Alert, 0, len(d.waiting))
for _, alert := range d.waiting {
waiting = append(waiting, *alert)
}
return waiting
}
// oldest returns the oldest alert waiting, nil when none waits.
func (d *destination) oldest() *Alert {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) == 0 {
return nil
}
return d.waiting[0]
}
// remove takes alert, which run has sent or given up, out of the queue,
// unless it has been dropped from it, or load has replaced the queue,
// since run took it. Only the oldest alert is ever dropped, so alert is
// the oldest if it is there at all.
func (d *destination) remove(alert *Alert) {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) > 0 && d.waiting[0] == alert {
d.waiting = slices.Delete(d.waiting, 0, 1)
}
}
// send posts alert to the destination, as message makes it, and returns
// an error unless the destination answers with a 2xx status: one that
// wraps errRefused for a 4xx status other than 408 and 429. No error
// names the destination's URL, whose path or query can carry a secret.
func (d *destination) send(ctx context.Context, alert *Alert) error {
body, header, err := d.message(alert)
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(ctx, sendTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, d.url.String(),
bytes.NewReader(body))
if err != nil {
return fmt.Errorf("make the request: %w", err)
}
maps.Copy(req.Header, header)
res, err := d.httpClient.Do(req)
if err != nil {
// The client's error names the URL: only what went wrong is kept.
if urlErr, ok := errors.AsType[*url.Error](err); ok {
return urlErr.Err
}
return err
}
defer func() {
_ = res.Body.Close()
}()
// Read, so that the connection can be used again.
_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, maxAnswerBytes))
switch status := res.StatusCode; {
case status >= http.StatusOK && status < http.StatusMultipleChoices:
return nil
case status >= http.StatusBadRequest && status < http.StatusInternalServerError &&
status != http.StatusRequestTimeout && status != http.StatusTooManyRequests:
return fmt.Errorf("%w %s", errRefused, res.Status)
default:
return fmt.Errorf("%w %s", errStatus, res.Status)
}
}
// webhookMessage returns alert as JSON, for the webhook, and the headers
// sent with it: WebhookHeaders, and its Content-Type.
func (q *Queue) webhookMessage(alert *Alert) ([]byte, http.Header, error) {
body, err := json.Marshal(alert)
if err != nil {
return nil, nil, fmt.Errorf("encode the alert: %w", err)
}
header := http.Header{}
maps.Copy(header, q.params.WebhookHeaders)
header.Set("Content-Type", "application/json")
return body, header, nil
}
// slackMessage returns alert as a message for a Slack incoming webhook,
// in JSON: its title in bold, then its text, with &, < and > escaped, as
// Slack asks, so that nothing in them is read as a link or a mention.
func slackMessage(alert *Alert) ([]byte, http.Header, error) {
escape := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;").Replace
body, err := json.Marshal(map[string]string{
"text": "*" + escape(title(alert)) + "*\n" + escape(text(alert)),
})
if err != nil {
return nil, nil, fmt.Errorf("encode the message: %w", err)
}
return body, http.Header{"Content-Type": {"application/json"}}, nil
}
// ntfyMessage returns alert's text, as the message published to ntfy,
// and the headers sent with it: its title, the priority and the tag of
// its event, and NtfyToken, unless it is empty, as a bearer token.
func (q *Queue) ntfyMessage(alert *Alert) ([]byte, http.Header, error) {
header := http.Header{
"Title": {title(alert)},
"Priority": {ntfyPriority(alert.Event)},
"Tags": {ntfyTag(alert.Event)},
}
if q.params.NtfyToken != "" {
header.Set("Authorization", "Bearer "+q.params.NtfyToken)
}
return []byte(text(alert)), header, nil
}
// ntfyPriority returns the priority an alert for event is published to
// ntfy with: high for an event the admin needs to look at.
func ntfyPriority(event string) string {
switch event {
case EventPermanentBan, EventAnomaly, EventSourceFailure, EventFileError:
return "high"
case EventReputationHit:
return "low"
default: // ban, waf_block and summary
return "default"
}
}
// ntfyTag returns the tag an alert for event is published to ntfy with,
// which ntfy shows as an emoji.
func ntfyTag(event string) string {
switch event {
case EventBan, EventPermanentBan:
return "no_entry"
case EventWAFBlock:
return "shield"
case EventAnomaly:
return "chart_with_upwards_trend"
case EventReputationHit:
return "label"
case EventSourceFailure, EventFileError:
return "warning"
default: // summary
return "bar_chart"
}
}
// title returns the title of alert in Slack and ntfy: the instance and
// the event.
func title(alert *Alert) string {
return alert.Instance + ": " + alert.Event
}
// text returns the text of alert in Slack and ntfy: its reason, then a
// line for each of its client, netblock and country, the file, source,
// error and mode its detail gives, and its suppressed repeats, that it
// has.
func text(alert *Alert) string {
lines := []string{alert.Reason}
if alert.Client.IsValid() {
lines = append(lines, "client: "+alert.Client.String())
}
if alert.Netblock.IsValid() {
lines = append(lines, "netblock: "+alert.Netblock.String())
}
if alert.Country != "" {
lines = append(lines, "country: "+alert.Country)
}
for _, name := range []string{"file", "source", "error", "mode"} {
value, _ := alert.Detail[name].(string)
if value != "" {
lines = append(lines, name+": "+value)
}
}
if alert.SuppressedRepeats > 0 {
lines = append(lines, fmt.Sprintf("suppressed repeats: %d", alert.SuppressedRepeats))
}
return strings.Join(lines, "\n")
}
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
package alerts
import "net/http"
// QueueSize is the most alerts that wait to be sent to a destination.
const QueueSize = queueSize
// SetTransport has q's requests to the destination name go through
// transport instead of the network.
func (q *Queue) SetTransport(name string, transport http.RoundTripper) {
for _, d := range q.destinations {
if d.name == name {
d.httpClient.Transport = transport
}
}
}
+286
View File
@@ -0,0 +1,286 @@
package bans_test
import (
"net/netip"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
func TestBanWithoutACauseIsAnAdmins(t *testing.T) {
t.Parallel()
netblock := netip.MustParsePrefix("203.0.113.0/24")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{{Netblock: netblock, Start: midnight()}})
if got := ledger.Bans(netblock)[0].Cause; got != bans.CauseAdmin {
t.Errorf("the ban's cause is %q, want admin", got)
}
}
func TestAdminsBansAreNeverDroppedAndDoNotCountTowardMaxBans(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
adminsOnly := netip.MustParsePrefix("198.51.100.0/24")
both := netip.MustParsePrefix("203.0.113.1/32")
second := netip.MustParsePrefix("203.0.113.2/32")
third := netip.MustParsePrefix("203.0.113.3/32")
// Seen longest ago, a netblock with two of an admin's bans alone, and
// then one with an admin's ban before a ban smallwebwaf made: the one
// ban counted toward MaxBans.
ledger.Load([]bans.Ban{
{Netblock: adminsOnly, Start: midnight().Add(-3 * time.Hour), Cause: bans.CauseAdmin},
{Netblock: adminsOnly, Start: midnight().Add(-2 * time.Hour), Cause: bans.CauseAdmin},
{Netblock: both, Start: midnight().Add(-time.Hour), Cause: bans.CauseAdmin},
{
Netblock: both,
Start: midnight(),
Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
},
})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 2})
// A new ban drops the ban smallwebwaf made, and only that one.
ledger.BanForLimit(second, midnight(), bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 1})
if ledger.Bans(both)[0].Cause != bans.CauseAdmin {
t.Errorf("%s kept %+v, want the admin's ban", both, ledger.Bans(both))
}
// And the next drops that one.
ledger.BanForLimit(third, midnight(), bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 0, third: 1})
}
func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{Limit: 1000, Window: "minute"})
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
bans.Notes{RuleID: "git-dir", Target: "path"})
for _, tc := range []struct{ got, want string }{
{limit.Reason, "requests per minute over the limit of 1000"},
{attack.Reason, "matched the rule git-dir"},
} {
if tc.got != tc.want {
t.Errorf("the reason is %q, want %q", tc.got, tc.want)
}
}
}
func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
t.Parallel()
// An hour's ban lifted ten minutes after it started.
netblock := netip.MustParsePrefix("203.0.113.9/32")
lifted := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
Lifted: midnight().Add(10 * time.Minute),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{lifted})
// While it would still last, it refuses nothing, and a limit broken
// bans for an hour, as a first broken limit does; the lifted ban is
// kept, and counted among the earlier bans.
now := midnight().Add(30 * time.Minute)
_, banned, _ := ledger.Check(netblock.Addr(), now)
if banned {
t.Error("the lifted ban refuses")
}
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
held := ledger.Bans(netblock)
if len(held) != 2 || held[0] != lifted {
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
}
}
func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
t.Parallel()
// A permanent ban for a clear sign of attack, lifted.
netblock := netip.MustParsePrefix("203.0.113.9/32")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{{
Netblock: netblock,
Start: midnight(),
Cause: bans.CauseAttack,
Lifted: midnight().Add(time.Hour),
}})
now := midnight().Add(2 * time.Hour)
_, banned, _ := ledger.Find(netblock.Addr(), now)
if banned {
t.Error("the lifted ban refuses")
}
active, permanent := ledger.Count(now)
if active != 0 || permanent != 0 {
t.Errorf("%d bans are active and %d permanent, want none", active, permanent)
}
// The next clear sign of attack bans for seven days, as a first does.
ban, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != 7*day {
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
}
}
func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
made, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{})
atStart := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
Start: midnight(),
}
// The bans read at the start were made before it.
ledger.Load([]bans.Ban{made, atStart})
if got := ledger.Made(bans.CauseAdmin); got != 0 {
t.Fatalf("%d bans made by an admin after the start's, want none", got)
}
// The admin keeps the ban smallwebwaf made, keeps the one read at the
// start, and adds one without a cause: that one alone is made.
kept := made
kept.Cause = bans.CauseAdmin
added := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.3/32"),
Start: midnight(),
}
ledger.LoadEdit([]bans.Ban{kept, atStart, added})
if ledger.Made(bans.CauseAdmin) != 1 || ledger.Made(bans.CauseLimit) != 1 {
t.Errorf("%d bans made by an admin and %d for a limit, want 1 of each",
ledger.Made(bans.CauseAdmin), ledger.Made(bans.CauseLimit))
}
}
func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
t.Parallel()
netblock := netip.MustParsePrefix("203.0.113.0/24")
ledger := bans.New(defaultRules())
// An hour's ban for a broken limit.
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// A minute later an admin bans the netblock for good, named by an
// address in it: that ban is made, and counts the other among the
// earlier bans.
now := midnight().Add(time.Minute)
want := bans.Ban{
Netblock: netblock,
Start: now,
Cause: bans.CauseAdmin,
Reason: "probes for logins",
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 1}},
}
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
"probes for logins")
if got != want {
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
}
wantChanged(t, ledger, true)
if made := ledger.Made(bans.CauseAdmin); made != 1 {
t.Errorf("%d bans made by an admin, want 1", made)
}
// It refuses once the ban for the limit has ended.
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
if !banned || ban != want {
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
ban, banned, want)
}
}
func TestLiftLiftsEveryActiveBanCoveringTheClient(t *testing.T) {
t.Parallel()
client := netip.MustParseAddr("203.0.113.9")
own := netip.MustParsePrefix("203.0.113.9/32")
wide := netip.MustParsePrefix("203.0.113.0/24")
other := netip.MustParsePrefix("203.0.113.10/32")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{
// Ended an hour ago.
{
Netblock: own, Start: midnight().Add(-2 * time.Hour),
Expires: midnight().Add(-time.Hour), Cause: bans.CauseLimit,
},
// Active, on the client's address and on its /24.
{
Netblock: own, Start: midnight(), Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
},
{Netblock: wide, Start: midnight(), Cause: bans.CauseAdmin},
// Another client's.
{Netblock: other, Start: midnight(), Cause: bans.CauseAdmin},
})
now := midnight().Add(time.Minute)
lifted := ledger.Lift(client, now)
if len(lifted) != 2 || lifted[0].Lifted != now || lifted[1].Lifted != now {
t.Errorf("lifted %+v, want the two active bans covering the client", lifted)
}
wantChanged(t, ledger, true)
if _, banned, _ := ledger.Check(client, now); banned {
t.Error("the client is still banned")
}
if _, banned, _ := ledger.Check(other.Addr(), now); !banned {
t.Error("the other client's ban was lifted")
}
// The lifted bans are kept, and the one that had ended is not lifted.
covering := ledger.Covering(client)
if len(covering) != 3 || covering[0].Netblock != wide ||
!covering[1].Lifted.IsZero() || covering[2].Lifted != now {
t.Errorf("the bans covering the client are %+v, want the /24's and both "+
"of its own, the earlier not lifted", covering)
}
// With none active, nothing is lifted or changed.
if lifted = ledger.Lift(client, now); len(lifted) != 0 {
t.Errorf("lifted %+v again", lifted)
}
wantChanged(t, ledger, false)
}
+637 -105
View File
@@ -1,9 +1,13 @@
// 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, with their notes, as the // netblocks of clients that break a rate limit or show a clear sign of
// "Bans" section of SPEC.md describes. The bans are kept in memory only. // attack, and 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 (
"fmt"
"math"
"net/netip" "net/netip"
"slices" "slices"
"strings" "strings"
@@ -13,6 +17,17 @@ import (
"github.com/hashicorp/golang-lru/v2/simplelru" "github.com/hashicorp/golang-lru/v2/simplelru"
) )
// The causes of bans.
const (
// CauseLimit is a ban smallwebwaf made for a broken limit.
CauseLimit = "limit"
// CauseAttack is a ban smallwebwaf made for a clear sign of attack.
CauseAttack = "attack"
// CauseAdmin is a ban an admin made, or one smallwebwaf made that an
// admin keeps. It is never dropped.
CauseAdmin = "admin"
)
// 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
// for a limit broken again within the repeat window lasts. // for a limit broken again within the repeat window lasts.
const repeatFactor = 3 const repeatFactor = 3
@@ -20,32 +35,44 @@ const repeatFactor = 3
// maxTextBytes is how much of each text in a ban's notes is kept. // maxTextBytes is how much of each text in a ban's notes is kept.
const maxTextBytes = 256 const maxTextBytes = 256
// Rules are how long a ban for a broken limit lasts, and how many bans // Rules are how long a ban lasts, and how many bans are held.
// are held.
type Rules struct { type Rules struct {
// LimitBanDuration is how long a first ban 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 netblock's last ban // LimitBanRepeatWindow is how soon after the end of the netblock's
// ended a broken limit counts as a repeat, which bans for // ban that ended last, other than one for a clear sign of attack, a
// repeatFactor times as long as that ban. // broken limit counts as a repeat, which bans for repeatFactor times as
// long as that ban.
LimitBanRepeatWindow time.Duration LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban; a ban that would be longer is // MaxBanDuration is the longest ban for a broken limit; one that would
// permanent instead. // be longer is permanent instead.
MaxBanDuration time.Duration MaxBanDuration time.Duration
// MaxBans is the most bans held, at least one. Past it, the earliest // AttackBanDuration is how long a first ban for a clear sign of attack
// ban of the netblock that has gone longest without a request is // lasts.
// dropped. AttackBanDuration time.Duration
// MaxBans is the most bans held whose cause is not CauseAdmin, at
// least one. Past it, the earliest such ban of the netblock that has
// gone longest without a request is dropped. Bans whose cause is
// CauseAdmin are held besides, and never dropped.
MaxBans int MaxBans int
} }
// Ban is a ban on a netblock for a broken limit, the only kind of ban // Ban is a ban on a netblock.
// smallwebwaf makes so far.
type Ban struct { type Ban struct {
Netblock netip.Prefix Netblock netip.Prefix
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
Notes Notes // Cause is CauseLimit, CauseAttack or CauseAdmin.
Cause string
// 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.
Reason string
// 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
// ban longer.
Lifted time.Time
Notes Notes
} }
// Permanent reports whether the ban never runs out. // Permanent reports whether the ban never runs out.
@@ -53,140 +80,323 @@ func (b Ban) Permanent() bool {
return b.Expires.IsZero() return b.Expires.IsZero()
} }
// ActiveAt reports whether the ban refuses requests at now. // ActiveAt reports whether the ban refuses requests at now: it has not
// been lifted, and has not run out.
func (b Ban) ActiveAt(now time.Time) bool { func (b Ban) ActiveAt(now time.Time) bool {
return b.Permanent() || now.Before(b.Expires) return b.Lifted.IsZero() && (b.Permanent() || now.Before(b.Expires))
} }
// Notes are what an admin needs to decide whether to lift a ban. // Notes are what an admin needs to decide whether to lift a ban. The
// JSON names are those of bans.json.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Notes struct { type Notes struct {
// Country is the client's country, when it was looked up. // Country is the client's country, when it was looked up.
Country string Country string `json:"country"`
// Limit, Window and Count are the limit that was broken, its window, // Limit, Window and Count are, for a ban for a broken limit, the limit
// "minute", "hour" or "day", and the count reached: the client's // that was broken, its window, "minute", "hour" or "day", and the
// requests in the window, the one that broke the limit included. // count reached: the client's requests in the window, the one that
// These are the requests that counted toward the ban, and the window // broke the limit included. These are the requests that counted
// is the time over which they came. // toward the ban, and the window is the time over which they came.
Limit int64 Limit int64 `json:"limit,omitempty"`
Window string Window string `json:"window,omitempty"`
Count float64 Count float64 `json:"count,omitempty"`
// Request is the request that broke the limit. // RuleID and Target are, for a ban for a clear sign of attack, the id
Request Request // of the rule file rule that matched, and its target.
// Refused is how many requests the ban has refused so far. RuleID string `json:"rule_id,omitempty"`
Refused int64 Target string `json:"target,omitempty"`
// EarlierBans is how many bans the netblock had before this one. // Request is the request that broke the limit, or that was the clear
EarlierBans int // sign of attack.
Request Request `json:"request"`
// 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
// far. Both go up with each request the ban refuses.
Requests int64 `json:"requests"`
Refused int64 `json:"refused"`
// EarlierBans is how many bans the netblock had before this one, by
// cause.
EarlierBans EarlierBans `json:"earlier_bans"`
}
// EarlierBans counts a netblock's bans before a ban, by cause.
type EarlierBans struct {
Limit int `json:"limit"`
Attack int `json:"attack"`
Admin int `json:"admin"`
} }
// 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.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Request struct { type Request struct {
Time time.Time Time time.Time `json:"time"`
Method string Method string `json:"method"`
Host string Host string `json:"host"`
// Path is the path with its query string. // Path is the path with its query string.
Path string Path string `json:"path"`
// Status is what the client was sent, 0 if nothing was. // Status is what the client was sent, 0 if nothing was.
Status int Status int `json:"status"`
UserAgent string UserAgent string `json:"user_agent"`
} }
// Ledger holds the bans. It is safe for concurrent use. // Ledger holds the bans. It is safe for concurrent use.
type Ledger struct { type Ledger struct {
rules Rules rules Rules
// changed receives a value when a ban is made, unless one is waiting
// already.
changed chan struct{}
mu sync.Mutex mu sync.Mutex
// netblocks holds each banned netblock's bans, oldest first. Each // netblocks holds each banned netblock's bans, oldest first. Check and
// request from a netblock makes it the most recently seen. // Find make each netblock they find the most recently seen.
netblocks *simplelru.LRU[netip.Prefix, *[]Ban] netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans. // held is how many bans netblocks holds whose cause is not CauseAdmin,
// at most rules.MaxBans.
held int held int
// made is how many bans have been made since the start, by cause: by
// the ledger, and by an admin, through BanForAdmin or in an edit of
// bans.json.
made map[string]int
// v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6
// netblocks that have been banned. Check looks for a ban at each of
// them, so that a ban read from bans.json refuses every client in its
// netblock even when it was made with another SWWAF_BAN_SCOPE_V4_PREFIX,
// or another length of an IPv6 client's netblock.
v4Lengths, v6Lengths []int
} }
// New returns a Ledger with no ban yet. // New returns a Ledger with no ban yet.
func New(rules Rules) *Ledger { func New(rules Rules) *Ledger {
// Every netblock held has a ban, so there are never more netblocks // The ledger drops bans itself, and never those whose cause is
// than rules.MaxBans, and the LRU never drops one itself. // CauseAdmin, however many there are, so the LRU has no limit of its
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](rules.MaxBans, nil) // own: it keeps the netblocks in the order they were last seen.
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](math.MaxInt, 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
} }
return &Ledger{rules: rules, netblocks: netblocks} return &Ledger{
rules: rules,
changed: make(chan struct{}, 1),
netblocks: netblocks,
made: map[string]int{},
}
} }
// Check is called for each request from netblock, at now. It reports // Changed receives a value after a ban is made, lifted or made permanent,
// whether a ban on netblock is active, and returns that ban, with the // so that bans.json can be written. Several changes before it is read
// request counted among those it refused. // leave one value.
func (l *Ledger) Check(netblock netip.Prefix, now time.Time) (Ban, bool) { func (l *Ledger) Changed() <-chan struct{} {
return l.changed
}
// Check is called for a request from client, at now. It reports whether
// a ban on a netblock client is in is active, and returns that ban, with
// the request counted among those it refused. A ban for a clear sign of
// attack is made permanent by the request: the netblock is malicious.
// The last result reports whether the request made the ban permanent.
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
bans, found := l.netblocks.Get(netblock) ban := l.active(client, now)
if !found { if ban == nil {
return Ban{}, false return Ban{}, false, false
} }
// A ban is made only once the one before has ended, so only the last ban.Notes.Requests++
// can be active. ban.Notes.Refused++
last := &(*bans)[len(*bans)-1]
if !last.ActiveAt(now) { madePermanent := ban.Cause == CauseAttack && !ban.Permanent()
return Ban{}, false if madePermanent {
ban.Expires = time.Time{}
l.markChanged()
} }
last.Notes.Refused++ return *ban, true, madePermanent
}
return *last, true // Find is Check without counting the request among those the ban
// refused, and without making the ban permanent: in observe mode a ban
// refuses nothing. The last result reports whether Check would have made
// the ban permanent.
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock()
defer l.mu.Unlock()
ban := l.active(client, now)
if ban == nil {
return Ban{}, false, false
}
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent()
}
// activeBan returns the ban in bans, a netblock's bans oldest first, that
// is active at now, or nil when none is. If several are, it returns the
// one that started last. Every ban is looked at, since a ban an admin adds
// to bans.json can start before the netblock's others and outlast them.
func activeBan(bans []Ban, now time.Time) *Ban {
for i := len(bans) - 1; i >= 0; i-- {
if bans[i].ActiveAt(now) {
return &bans[i]
}
}
return nil
} }
// BanForLimit bans netblock at now for a broken limit, with notes, and // BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban. A first ban lasts LimitBanDuration. A ban made within // returns the ban, and true. A first ban lasts LimitBanDuration. A ban
// LimitBanRepeatWindow after the netblock's last ban ended lasts // 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
// 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 and no other is made. The ledger fills in the notes' Refused // returned with false, and no other is made. The ledger fills in the
// and EarlierBans itself. // notes' Refused and EarlierBans itself, and gives the ban the reason
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban { // "requests per <Window> over the limit of <Limit>", from the notes.
func (l *Ledger) BanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
}
// WouldBanForLimit returns what BanForLimit would, without making the ban:
// what observe mode would have done.
func (l *Ledger) WouldBanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
}
// BanForAttack bans netblock at now for a clear sign of attack, with
// notes, and returns the ban, and whether it made it, as BanForLimit
// does. A first ban lasts AttackBanDuration; once the netblock has had
// one that was not lifted, the next is permanent. Its reason is "matched
// the rule <RuleID>".
func (l *Ledger) BanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true)
}
// WouldBanForAttack returns what BanForAttack would, without making the
// ban: what observe mode would have done.
func (l *Ledger) WouldBanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false)
}
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit
// or CauseAttack, made at now would be permanent, as BanForLimit or
// BanForAttack would make it. It works out nothing else of the ban.
func (l *Ledger) WouldBePermanent(
netblock netip.Prefix, now time.Time, cause string,
) bool {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
var last *Ban var held []Ban
if bans, found := l.netblocks.Peek(netblock); found {
bans, found := l.netblocks.Get(netblock) held = *bans
if found {
last = &(*bans)[len(*bans)-1]
if last.ActiveAt(now) {
return *last
}
notes.EarlierBans = last.Notes.EarlierBans + 1
} }
notes.Request = notes.Request.cut() if cause == CauseAttack {
return l.attackExpiry(held, now).IsZero()
}
return l.limitExpiry(held, now).IsZero()
}
// limitReason is the reason of a ban for a broken limit, with notes.
func limitReason(notes Notes) string {
return fmt.Sprintf("requests per %s over the limit of %d", notes.Window, notes.Limit)
}
// attackReason is the reason of a ban for a clear sign of attack, with
// notes.
func attackReason(notes Notes) string {
return "matched the rule " + notes.RuleID
}
// BanForAdmin bans netblock at now for an admin, with reason, until
// expires, or for good when expires is zero, and returns the ban, whose
// cause is CauseAdmin. Unlike BanForLimit and BanForAttack, it makes the
// ban even while another on netblock is active, since the admin asked
// for this one. The ledger fills in the notes' EarlierBans, and counts
// the ban among those made.
func (l *Ledger) BanForAdmin(
netblock netip.Prefix, now, expires time.Time, reason string,
) Ban {
l.mu.Lock()
defer l.mu.Unlock()
ban := Ban{ ban := Ban{
Netblock: netblock, Netblock: netblock.Masked(), Start: now, Expires: expires, Cause: CauseAdmin,
Start: now, Reason: reason,
Expires: l.expiry(last, now),
Notes: notes,
} }
if l.held == l.rules.MaxBans { held, found := l.netblocks.Get(ban.Netblock)
l.dropOne() if found {
ban.Notes.EarlierBans = earlierBans(*held)
} }
// dropOne can have dropped netblock's last ban, and netblock with it. l.add(ban)
bans, found = l.netblocks.Peek(netblock) l.made[CauseAdmin]++
if !found { l.markChanged()
bans = &[]Ban{}
l.netblocks.Add(netblock, bans)
}
*bans = append(*bans, ban)
l.held++
return ban return ban
} }
// Lift lifts, at now, every ban active then on a netblock client is in,
// as an admin does, and returns those bans. A lifted ban is kept, refuses
// nothing, and does not make the netblock's next ban longer.
func (l *Ledger) Lift(client netip.Addr, now time.Time) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
var lifted []Ban
for _, bans := range l.covering(client) {
for i := range *bans {
ban := &(*bans)[i]
if ban.ActiveAt(now) {
ban.Lifted = now
lifted = append(lifted, *ban)
}
}
}
if len(lifted) > 0 {
l.markChanged()
}
return lifted
}
// Covering returns every ban held on a netblock client is in, active or
// not, sorted by netblock, and each netblock's bans oldest first. It is
// not a request from client, and leaves when the netblocks were last seen
// unchanged.
func (l *Ledger) Covering(client netip.Addr) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
var held []Ban
for _, bans := range l.covering(client) {
held = append(held, *bans...)
}
slices.SortStableFunc(held, func(a, b Ban) int {
return a.Netblock.Compare(b.Netblock)
})
return held
}
// Bans returns the bans held on netblock, oldest first. It is not a // Bans returns the bans held on netblock, oldest first. It is not a
// request from netblock, and leaves when it was last seen unchanged. // request from netblock, and leaves when it was last seen unchanged.
func (l *Ledger) Bans(netblock netip.Prefix) []Ban { func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
@@ -201,12 +411,298 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans) return slices.Clone(*bans)
} }
// expiry returns when a ban for a broken limit made at now ends, or zero // Made returns how many bans for cause have been made since the start:
// when it is permanent. last is the netblock's last ban, which has ended, // for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
// or nil when it has none. // admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time { // them. The bans read from bans.json at the start are not among them.
func (l *Ledger) Made(cause string) int {
l.mu.Lock()
defer l.mu.Unlock()
return l.made[cause]
}
// Count returns how many of the bans held are active at now, and how many
// of those are permanent. A lifted ban is neither.
func (l *Ledger) Count(now time.Time) (int, int) {
l.mu.Lock()
defer l.mu.Unlock()
active, permanent := 0, 0
for _, bans := range l.netblocks.Values() {
for _, ban := range *bans {
if !ban.ActiveAt(now) {
continue
}
active++
if ban.Permanent() {
permanent++
}
}
}
return active, permanent
}
// Snapshot returns every ban held, sorted by netblock, and each
// netblock's bans oldest first, as bans.json lists them.
func (l *Ledger) Snapshot() []Ban {
l.mu.Lock()
defer l.mu.Unlock()
held := make([]Ban, 0, l.held)
for _, bans := range l.netblocks.Values() {
held = append(held, *bans...)
}
slices.SortStableFunc(held, func(a, b Ban) int {
return a.Netblock.Compare(b.Netblock)
})
return held
}
// Load puts bans read from bans.json at the start into the ledger, in
// place of the bans it holds, in the order they started, so that a
// netblock whose last ban started latest counts as the most recently
// seen. A ban without a cause is an admin's, and gets CauseAdmin. Each
// netblock is masked to its length, so that 203.0.113.9/24 is
// 203.0.113.0/24, and each text in the notes is cut to 256 bytes. Past
// MaxBans the earliest bans whose cause is not CauseAdmin are dropped, as
// when they are made.
func (l *Ledger) Load(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
l.load(bans)
}
// LoadEdit is Load for an admin's edit of bans.json, taken in while
// smallwebwaf runs. Each ban in it whose cause is CauseAdmin, and which
// the ledger did not hold, with the same netblock and start, is one the
// admin made, and is counted among the bans made.
func (l *Ledger) LoadEdit(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
l.made[CauseAdmin] += l.load(bans)
}
// load does what Load describes, and returns how many of bans are bans
// whose cause is CauseAdmin that the ledger did not hold before.
func (l *Ledger) load(bans []Ban) int {
bans = slices.Clone(bans)
added := 0
for i := range bans {
ban := &bans[i]
ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut()
if ban.Cause == "" {
ban.Cause = CauseAdmin
}
if ban.Cause == CauseAdmin && !l.holds(ban.Netblock, ban.Start) {
added++
}
}
slices.SortStableFunc(bans, func(a, b Ban) int {
return a.Start.Compare(b.Start)
})
l.netblocks.Purge()
l.held = 0
l.v4Lengths, l.v6Lengths = nil, nil
for _, ban := range bans {
l.add(ban)
}
return added
}
// holds reports whether the ledger holds a ban on netblock that started
// at start.
func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
bans, found := l.netblocks.Peek(netblock)
return found && slices.ContainsFunc(*bans, func(ban Ban) bool {
return ban.Start.Equal(start)
})
}
// ban bans netblock at now for cause, with reason and notes, as
// BanForLimit and BanForAttack describe, and returns the ban, and whether
// it made it. Unless keep is true, the ban is not made, only returned: it
// is the ban that would have been made.
func (l *Ledger) ban(
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool,
) (Ban, bool) {
l.mu.Lock()
defer l.mu.Unlock()
// held are the netblock's bans, none of them active.
var held []Ban
bans, found := l.netblocks.Get(netblock)
if found {
active := activeBan(*bans, now)
if active != nil {
return *active, false
}
held = *bans
notes.EarlierBans = earlierBans(held)
}
notes.Request = notes.Request.cut()
ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes}
if cause == CauseAttack {
ban.Expires = l.attackExpiry(held, now)
} else {
ban.Expires = l.limitExpiry(held, now)
}
if !keep {
return ban, true
}
l.add(ban)
l.made[cause]++
l.markChanged()
return ban, true
}
// earlierBans returns how many bans a netblock with the bans held, oldest
// first, has had, by cause: the first ban held counts the bans the
// netblock had before that one, since dropped to make room, and each ban
// held adds one.
func earlierBans(held []Ban) EarlierBans {
earlier := held[0].Notes.EarlierBans
for _, ban := range held {
switch ban.Cause {
case CauseLimit:
earlier.Limit++
case CauseAttack:
earlier.Attack++
case CauseAdmin:
earlier.Admin++
}
}
return earlier
}
// markChanged has Changed receive a value, unless one is waiting already.
func (l *Ledger) markChanged() {
select {
case l.changed <- struct{}{}:
default: // a value is waiting already
}
}
// active returns the ban active at now on a netblock client is in, or
// nil.
func (l *Ledger) active(client netip.Addr, now time.Time) *Ban {
lengths := l.v6Lengths
if client.Is4() {
lengths = l.v4Lengths
}
for _, length := range lengths {
bans, found := l.netblocks.Get(netip.PrefixFrom(client, length).Masked())
if !found {
continue
}
ban := activeBan(*bans, now)
if ban != nil {
return ban
}
}
return nil
}
// covering returns the bans of each netblock held that client is in,
// leaving when the netblocks were last seen unchanged.
func (l *Ledger) covering(client netip.Addr) []*[]Ban {
lengths := l.v6Lengths
if client.Is4() {
lengths = l.v4Lengths
}
var found []*[]Ban
for _, length := range lengths {
bans, ok := l.netblocks.Peek(netip.PrefixFrom(client, length).Masked())
if ok {
found = append(found, bans)
}
}
return found
}
// add adds ban to its netblock's bans, after the last, and makes its
// netblock the most recently seen. With MaxBans held, it drops one first,
// unless ban's cause is CauseAdmin, which does not count toward MaxBans.
func (l *Ledger) add(ban Ban) {
counted := ban.Cause != CauseAdmin
if counted && l.held == l.rules.MaxBans {
l.dropOne()
}
// dropOne can have dropped the netblock's last ban, and the netblock
// with it.
bans, found := l.netblocks.Get(ban.Netblock)
if !found {
bans = &[]Ban{}
l.netblocks.Add(ban.Netblock, bans)
}
*bans = append(*bans, ban)
if counted {
l.held++
}
lengths := &l.v6Lengths
if ban.Netblock.Addr().Is4() {
lengths = &l.v4Lengths
}
if !slices.Contains(*lengths, ban.Netblock.Bits()) {
*lengths = append(*lengths, ban.Netblock.Bits())
}
}
// 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
// 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
// admin adds to bans.json can start after 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 {
length := l.rules.LimitBanDuration length := l.rules.LimitBanDuration
var last *Ban
for i, ban := range held {
if ban.Cause != CauseAttack && ban.Lifted.IsZero() &&
(last == nil || ban.Expires.After(last.Expires)) {
last = &held[i]
}
}
if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow { if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow {
lastLength := last.Expires.Sub(last.Start) lastLength := last.Expires.Sub(last.Start)
// This is repeatFactor * lastLength > MaxBanDuration, written so // This is repeatFactor * lastLength > MaxBanDuration, written so
@@ -225,17 +721,53 @@ func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
return now.Add(length) return now.Add(length)
} }
// dropOne drops the earliest ban of the netblock that has gone longest // attackExpiry returns when a ban for a clear sign of attack made at now
// without a request, and the netblock with it if that was its only ban. // ends. held are the netblock's bans, none of them active: if one of them
func (l *Ledger) dropOne() { // is for a clear sign of attack too, and was not lifted, the new ban is
netblock, bans, _ := l.netblocks.GetOldest() // permanent, and its end zero; otherwise it ends AttackBanDuration later.
if len(*bans) == 1 { func (l *Ledger) attackExpiry(held []Ban, now time.Time) time.Time {
l.netblocks.Remove(netblock) for _, ban := range held {
} else { if ban.Cause == CauseAttack && ban.Lifted.IsZero() {
*bans = slices.Delete(*bans, 0, 1) return time.Time{}
}
} }
l.held-- return now.Add(l.rules.AttackBanDuration)
}
// dropOne drops the earliest ban whose cause is not CauseAdmin of the
// netblock that has gone longest without a request, of those that hold
// such a ban, and the netblock with it if that was its only ban. It is
// called with at least one such ban held. It looks at each netblock once
// at most, and drops nothing when none holds such a ban.
func (l *Ledger) dropOne() {
for range l.netblocks.Len() {
netblock, bans, _ := l.netblocks.GetOldest()
i := slices.IndexFunc(*bans, func(ban Ban) bool {
return ban.Cause != CauseAdmin
})
if i < 0 {
// Its bans are all an admin's, and never dropped. Get makes
// it the most recently seen, so that the next netblock is
// looked at; when it was seen matters only for dropping a
// ban, and a ban added to it makes it the most recently seen
// anyway.
l.netblocks.Get(netblock)
continue
}
if len(*bans) == 1 {
l.netblocks.Remove(netblock)
} else {
*bans = slices.Delete(*bans, i, i+1)
}
l.held--
return
}
} }
// cut returns r with each text cut to maxTextBytes and copied, so that // cut returns r with each text cut to maxTextBytes and copied, so that
+264 -28
View File
@@ -21,11 +21,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and // Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
// 81 hours. // 81 hours.
for i, hours := range []int{1, 3, 9, 27, 81} { for i, hours := range []int{1, 3, 9, 27, 81} {
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
length := time.Duration(hours) * time.Hour length := time.Duration(hours) * time.Hour
if !ban.Expires.Equal(now.Add(length)) || ban.Notes.EarlierBans != i { if !ban.Expires.Equal(now.Add(length)) ||
t.Fatalf("ban %d lasts %s with %d earlier bans, want %d hours and %d", ban.Notes.EarlierBans != (bans.EarlierBans{Limit: i}) {
t.Fatalf("ban %d lasts %s with earlier bans %+v, want %d hours and %d for a limit",
i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i) i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i)
} }
@@ -34,12 +35,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// The sixth would last 243 hours, more than seven days: it is // The sixth would last 243 hours, more than seven days: it is
// permanent, and never ends. // permanent, and never ends.
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires) t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
} }
_, banned := ledger.Check(netblock, now.Add(100*365*day)) _, banned, _ := ledger.Check(netblock.Addr(), now.Add(100*365*day))
if !banned { if !banned {
t.Error("a permanent ban ended") t.Error("a permanent ban ended")
} }
@@ -63,11 +64,12 @@ func TestRepeatWindowRunsOut(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{}) second, _ := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
if second.Expires.Sub(second.Start) != tc.want || second.Notes.EarlierBans != 1 { if second.Expires.Sub(second.Start) != tc.want ||
t.Errorf("second ban lasts %s with %d earlier bans, want %s and 1", second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("second ban lasts %s with earlier bans %+v, want %s and 1 for a limit",
second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want) second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want)
} }
}) })
@@ -81,7 +83,7 @@ func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
ledger := bans.New(rules) ledger := bans.New(rules)
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), ban, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{}) bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires) t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
@@ -101,7 +103,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now := midnight() now := midnight()
for i := range 14 { for i := range 14 {
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Expires.After(ban.Start) { if !ban.Expires.After(ban.Start) {
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires) t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
} }
@@ -109,7 +111,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now = ban.Expires now = ban.Expires
} }
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires) t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
} }
@@ -121,12 +123,22 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) if !made {
t.Error("the first ban was not made")
}
if again != first || len(ledger.Bans(netblock)) != 1 { again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
again, len(ledger.Bans(netblock)), first) if made || again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
}
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first {
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
again, made, first)
} }
} }
@@ -135,28 +147,52 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
for range 3 { for range 3 {
got, banned := ledger.Check(netblock, ban.Expires.Add(-time.Nanosecond)) got, banned, _ := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got.Start != ban.Start { if !banned || got.Start != ban.Start {
t.Fatalf("check during the ban gives %+v and %t", got, banned) t.Fatalf("check during the ban gives %+v and %t", got, banned)
} }
} }
_, banned := ledger.Check(netip.MustParsePrefix("203.0.113.10/32"), midnight()) _, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
if banned { if banned {
t.Error("another netblock is banned") t.Error("another netblock is banned")
} }
_, banned = ledger.Check(netblock, ban.Expires) _, banned, _ = ledger.Check(netblock.Addr(), ban.Expires)
if banned { if banned {
t.Error("the ban did not end") t.Error("the ban did not end")
} }
refused := ledger.Bans(netblock)[0].Notes.Refused // The netblock's requests went from 5 to 8 with the three refused.
if refused != 3 { notes := ledger.Bans(netblock)[0].Notes
t.Errorf("the notes count %d refused requests, want 3", refused) if notes.Refused != 3 || notes.Requests != 8 {
t.Errorf("the notes count %d refused requests of %d, want 3 of 8",
notes.Refused, notes.Requests)
}
}
func TestFindCountsNothing(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
got, banned, _ := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got != ban {
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
}
_, banned, _ = ledger.Find(netblock.Addr(), ban.Expires)
if banned {
t.Error("the ban did not end")
}
if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes {
t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes)
} }
} }
@@ -172,13 +208,13 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
d := netip.MustParsePrefix("2001:db8::/64") d := netip.MustParsePrefix("2001:db8::/64")
now := midnight() now := midnight()
first := ledger.BanForLimit(a, now, bans.Notes{}) first, _ := ledger.BanForLimit(a, now, bans.Notes{})
ledger.BanForLimit(b, now, bans.Notes{}) ledger.BanForLimit(b, now, bans.Notes{})
ledger.BanForLimit(c, now, bans.Notes{}) ledger.BanForLimit(c, now, bans.Notes{})
// A request from a makes b the netblock seen longest ago, and its ban // A request from a makes b the netblock seen longest ago, and its ban
// goes to make room for d's. // goes to make room for d's.
ledger.Check(a, now) ledger.Check(a.Addr(), now)
ledger.BanForLimit(d, now, bans.Notes{}) ledger.BanForLimit(d, now, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1}) wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
@@ -188,7 +224,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
// With d seen since, a is seen longest ago, and its earlier ban goes // With d seen since, a is seen longest ago, and its earlier ban goes
// first. // first.
ledger.Check(d, first.Expires) ledger.Check(d.Addr(), first.Expires)
ledger.BanForLimit(b, first.Expires, bans.Notes{}) ledger.BanForLimit(b, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1}) wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
@@ -197,6 +233,205 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
} }
} }
func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
t.Parallel()
// With room for one ban, the netblock's ended ban goes to make room for
// its new one, whose notes still count it.
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
held := ledger.Bans(netblock)
if len(held) != 1 || held[0] != second ||
held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the ledger holds %+v, want only the second ban, "+
"with 1 earlier ban for a limit", held)
}
}
func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
notes := bans.Notes{RuleID: "env-file", Target: "path"}
ban, _ := ledger.BanForAttack(netblock, midnight(), notes)
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
ledger.Made(bans.CauseLimit) != 0 {
t.Fatalf("the ban is %+v, with %d made for an attack and %d for a limit, "+
"want one for an attack, of seven days", ban,
ledger.Made(bans.CauseAttack), ledger.Made(bans.CauseLimit))
}
wantChanged(t, ledger, true)
// In observe mode the ban refuses nothing, and stays as it is, while
// Find tells that the request would have made it permanent.
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent {
t.Fatalf("a request found under the ban left it %+v, would have made it "+
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
}
wantChanged(t, ledger, false)
// A request it refuses makes it permanent, says so, and makes
// bans.json due.
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
if !madePermanent || !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
t.Fatalf("after a request during the ban, it is %+v, made permanent %t, "+
"want it made permanent", got, madePermanent)
}
wantChanged(t, ledger, true)
// The next request finds it permanent already.
_, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
if !banned || madePermanent {
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
"want banned by the permanent ban", banned, madePermanent)
}
}
func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
// A ban for a broken limit before does not count.
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second, _ := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
if second.Expires.Sub(second.Start) != 7*day {
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
second.Expires.Sub(second.Start))
}
// Once that has run out without a request, the netblock is served, and
// its next clear sign of attack bans it for good.
_, banned, _ := ledger.Check(netblock.Addr(), second.Expires)
if banned {
t.Fatal("the ban did not end")
}
// Its notes show the earlier ban for an attack that makes it permanent,
// beside the one for a limit.
third, _ := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
if !third.Permanent() ||
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
"with 1 earlier ban for a limit and 1 for an attack", third)
}
}
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// While the first ban lasts, none would be made.
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
if would || during != first {
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
would, during, first)
}
// As it ends, a clear sign of attack would ban for seven days, and a
// limit broken again for three hours, but neither is made.
limitNotes := bans.Notes{Limit: 1, Window: "minute"}
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
bans.Notes{RuleID: "git-dir"})
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
if !wouldAttack || !attack.Expires.Equal(first.Expires.Add(7*day)) ||
attack.Reason != "matched the rule git-dir" || !wouldLimit ||
!limit.Expires.Equal(first.Expires.Add(3*time.Hour)) ||
limit.Reason != "requests per minute over the limit of 1" {
t.Errorf("would ban with %+v and %+v, want seven days for the attack and "+
"three hours for the limit", attack, limit)
}
if len(ledger.Bans(netblock)) != 1 || ledger.Made(bans.CauseLimit) != 1 ||
ledger.Made(bans.CauseAttack) != 0 {
t.Errorf("the ledger holds %+v, want the first ban alone", ledger.Bans(netblock))
}
wantChanged(t, ledger, false)
// The ban made is the one that would have been.
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
if made != limit {
t.Errorf("the ban made is %+v, want %+v", made, limit)
}
}
func TestWouldBePermanentAnswersAsTheBanWouldBeMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
// Five bans for a limit in a row, of 1, 3, 9, 27 and 81 hours, are not
// permanent. The sixth, of 243 hours, would be, while a first ban for
// an attack would not.
for i := range 5 {
if ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
t.Fatalf("ban %d for a limit would be permanent", i+1)
}
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
now = ban.Expires
}
if !ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
t.Error("the sixth ban for a limit would not be permanent")
}
if ledger.WouldBePermanent(netblock, now, bans.CauseAttack) {
t.Error("a first ban for an attack would be permanent")
}
// Once a first ban for an attack has ended, the next would be permanent.
attack, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
if !ledger.WouldBePermanent(netblock, attack.Expires, bans.CauseAttack) {
t.Error("a second ban for an attack would not be permanent")
}
}
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
// Three times the seven days would be permanent; a limit broken as the
// ban for an attack ends bans for an hour, as a first broken limit does.
attack, _ := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
limit, _ := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
}
// And a request during the ban for a limit leaves it as it is.
got, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start)
if got.Permanent() || madePermanent {
t.Error("a request during a ban for a limit made it permanent")
}
}
func TestRequestTextsAreCutTo256Bytes(t *testing.T) { func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
t.Parallel() t.Parallel()
@@ -207,7 +442,7 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long, Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
} }
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request}) ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
cut := long[:256] cut := long[:256]
want := bans.Request{ want := bans.Request{
@@ -225,6 +460,7 @@ func defaultRules() bans.Rules {
LimitBanDuration: time.Hour, LimitBanDuration: time.Hour,
LimitBanRepeatWindow: day, LimitBanRepeatWindow: day,
MaxBanDuration: 7 * day, MaxBanDuration: 7 * day,
AttackBanDuration: 7 * day,
MaxBans: 5000, MaxBans: 5000,
} }
} }
+314
View File
@@ -0,0 +1,314 @@
package bans_test
import (
"net/netip"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
func TestChangedAfterABanIsMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
wantChanged(t, ledger, false)
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// A limit broken during the ban makes no other, and a refusal changes
// only the counts in the notes, which wait for the interval's write.
ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
ledger.Check(netblock.Addr(), midnight().Add(time.Minute))
wantChanged(t, ledger, false)
// Two bans before the value is read leave one.
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{})
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.11/32"), midnight(), bans.Notes{})
wantChanged(t, ledger, true)
wantChanged(t, ledger, false)
}
func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
v6 := netip.MustParsePrefix("2001:db8::/64")
high := netip.MustParsePrefix("203.0.113.10/32")
low := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(v6, midnight(), bans.Notes{})
ledger.BanForLimit(high, midnight(), bans.Notes{})
ledger.BanForLimit(low, midnight(), bans.Notes{})
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
snapshot := ledger.Snapshot()
got := make([]string, 0, len(snapshot))
for _, ban := range snapshot {
got = append(got, ban.Netblock.String()+" "+ban.Start.Format(time.Kitchen))
}
want := []string{
"203.0.113.9/32 12:00AM", "203.0.113.10/32 12:00AM",
"2001:db8::/64 12:00AM", "2001:db8::/64 1:00AM",
}
if !slices.Equal(got, want) {
t.Errorf("snapshot %v, want %v", got, want)
}
}
func TestLoadedBansCarryOn(t *testing.T) {
t.Parallel()
before := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban, _ := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
// Loaded into a new ledger, as across a restart, the ban still refuses
// while it lasts, and once it has ended a broken limit bans for three
// times as long, with the loaded ban counted among the earlier ones.
after := bans.New(defaultRules())
after.Load(before.Snapshot())
_, banned, _ := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
if !banned {
t.Error("the loaded ban does not refuse")
}
again, _ := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
if again.Expires.Sub(again.Start) != 3*time.Hour ||
again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
again.Expires.Sub(again.Start), again.Notes.EarlierBans)
}
}
func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
t.Parallel()
// Two entries as an admin might write them, with addresses not masked
// to their lengths, the IPv6 one shorter than the /64 an IPv6 client's
// ban covers, beside a ban the ledger makes on one IPv4 address.
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.9/24"), Start: midnight()},
{Netblock: netip.MustParsePrefix("2001:db8::1/48"), Start: midnight()},
})
ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), bans.Notes{})
for client, want := range map[string]bool{
"203.0.113.0": true,
"203.0.113.200": true,
"203.0.114.1": false,
"2001:db8:0:5::1": true,
"2001:db8:1::1": false,
"198.51.100.7": true,
"198.51.100.8": false,
} {
_, banned, _ := ledger.Check(netip.MustParseAddr(client), midnight())
if banned != want {
t.Errorf("%s is refused: %t, want %t", client, banned, want)
}
}
// The loaded netblocks are written back masked.
snapshot := ledger.Snapshot()
got := make([]string, 0, len(snapshot))
for _, ban := range snapshot {
got = append(got, ban.Netblock.String())
}
want := []string{"198.51.100.7/32", "203.0.113.0/24", "2001:db8::/48"}
if !slices.Equal(got, want) {
t.Errorf("the ledger holds bans on %v, want %v", got, want)
}
}
func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
t.Parallel()
// As when an admin adds a permanent ban to bans.json with a start
// before that of the netblock's ban that has ended.
netblock := netip.MustParsePrefix("203.0.113.0/24")
permanent := bans.Ban{Netblock: netblock, Start: midnight().Add(-time.Hour)}
ended := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(time.Hour),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{permanent, ended})
now := midnight().Add(2 * time.Hour)
client := netip.MustParseAddr("203.0.113.9")
ban, banned, _ := ledger.Find(client, now)
if !banned || !ban.Permanent() {
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
}
ban, banned, _ = ledger.Check(client, now)
if !banned || !ban.Permanent() {
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
banned, ban)
}
// A limit broken now makes no shorter ban over the permanent one.
ban, _ = ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
}
}
func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
t.Parallel()
// A 9-hour ban smallwebwaf made, the third in a row, and an admin's
// 1-hour ban added to bans.json over it, with no cause and no notes.
netblock := netip.MustParsePrefix("203.0.113.9/32")
nineHours := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(9 * time.Hour),
Cause: bans.CauseLimit,
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 2}},
}
admins := bans.Ban{
Netblock: netblock,
Start: midnight().Add(time.Hour),
Expires: midnight().Add(2 * time.Hour),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{nineHours, admins})
// Once both have ended, a limit broken within the repeat window bans
// for three times the 9 hours, and the notes count the two bans
// before the 9-hour one and it, for a limit, and the admin's.
ban, _ := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
"want 27h, 3 for a limit and 1 an admin's",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
}
func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
t.Parallel()
// bans.json lists the bans by netblock, not in the order they began.
later := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.1/32"),
Start: midnight(),
Cause: bans.CauseLimit,
}
earlier := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
Start: midnight().Add(-time.Hour),
Cause: bans.CauseLimit,
}
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
ledger.Load([]bans.Ban{later, earlier})
held := ledger.Snapshot()
if len(held) != 1 || held[0] != later {
t.Errorf("the ledger holds %+v, want only the ban that began later", held)
}
}
func TestLoadReplacesTheBansHeld(t *testing.T) {
t.Parallel()
// Room for three bans, so that the second load, were it added to the
// two bans held, would drop none of them to make room.
rules := defaultRules()
rules.MaxBans = 3
ledger := bans.New(rules)
kept := bans.Ban{
Netblock: netip.MustParsePrefix("2001:db8::/64"),
Start: midnight(),
Cause: bans.CauseLimit,
}
ledger.Load([]bans.Ban{
{
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Start: midnight(),
Cause: bans.CauseLimit,
},
kept,
})
// Loaded again without the first ban, as when an admin's edit of
// bans.json is taken in, that ban is lifted.
ledger.Load([]bans.Ban{kept})
_, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
if banned {
t.Error("a ban left out of the second load still refuses")
}
// The ledger holds one ban, so it makes two more without dropping any.
first, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
bans.Notes{})
second, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
bans.Notes{})
want := []bans.Ban{first, second, kept}
if got := ledger.Snapshot(); !slices.Equal(got, want) {
t.Errorf("the ledger holds %+v, want %+v", got, want)
}
}
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
t.Parallel()
long := strings.Repeat("a", 300)
ban := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.9/32"),
Start: midnight(),
Notes: bans.Notes{Request: bans.Request{
Method: long, Host: long, Path: long, UserAgent: long,
}},
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{ban})
cut := long[:256]
want := bans.Request{Method: cut, Host: cut, Path: cut, UserAgent: cut}
got := ledger.Snapshot()[0].Notes.Request
if got != want {
t.Errorf("the notes keep %+v, want each text cut to 256 bytes", got)
}
}
// wantChanged checks whether the ledger's Changed has a value to read.
func wantChanged(t *testing.T, ledger *bans.Ledger, want bool) {
t.Helper()
got := false
select {
case <-ledger.Changed():
got = true
default:
}
if got != want {
t.Errorf("Changed has a value: %t, want %t", got, want)
}
}
+701 -17
View File
@@ -1,9 +1,11 @@
// Package config reads smallwebwaf's settings. Every setting is an // Package config reads smallwebwaf's settings. Every setting is an
// environment variable whose name starts with SWWAF_, every setting has a // environment variable whose name starts with SWWAF_, or a file such a
// default, and this package is the one place they are read. // variable names, every setting has a default, and this package is the
// one place they are read.
package config package config
import ( import (
"crypto/x509"
"errors" "errors"
"fmt" "fmt"
"log/slog" "log/slog"
@@ -12,10 +14,17 @@ import (
"net/http" "net/http"
"net/netip" "net/netip"
"net/url" "net/url"
"os"
"path/filepath"
"slices" "slices"
"strconv" "strconv"
"strings" "strings"
"time" "time"
"unicode"
"unicode/utf8"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
) )
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero // Config is smallwebwaf's settings. A timeout, size or rate limit of zero
@@ -25,12 +34,28 @@ type Config struct {
ListenAddr string ListenAddr string
// UpstreamURL is the app (SWWAF_UPSTREAM_URL). // UpstreamURL is the app (SWWAF_UPSTREAM_URL).
UpstreamURL *url.URL UpstreamURL *url.URL
// InstanceName is the name each request log line gives as instance
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets
// to the first 12 characters of the container's id.
InstanceName string
// Observe is true in observe mode, when SWWAF_MODE is observe rather
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country
// lists, a rate limit or a rule would refuse is passed to the app
// instead, and no ban is made.
Observe bool
// TrustedProxies are the netblocks whose X-Forwarded-For is // TrustedProxies are the netblocks whose X-Forwarded-For is
// believed (SWWAF_TRUSTED_PROXIES). // believed (SWWAF_TRUSTED_PROXIES).
TrustedProxies []netip.Prefix TrustedProxies []netip.Prefix
// ClientRequestTimeout bounds reading the whole request from the // ClientRequestTimeout bounds reading the whole request from the
// client (SWWAF_CLIENT_REQUEST_TIMEOUT). // client (SWWAF_CLIENT_REQUEST_TIMEOUT).
ClientRequestTimeout time.Duration ClientRequestTimeout time.Duration
// ClientRequestHeaderMaxBytes is the largest request line and headers
// a client may send (SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES). It is
// never off, and always more than 4K.
ClientRequestHeaderMaxBytes int64
// ClientIdleTimeout bounds how long a kept-open client connection
// may wait for its next request (SWWAF_CLIENT_IDLE_TIMEOUT).
ClientIdleTimeout time.Duration
// ClientResponseTimeout bounds writing the whole response to the // ClientResponseTimeout bounds writing the whole response to the
// client (SWWAF_CLIENT_RESPONSE_TIMEOUT). // client (SWWAF_CLIENT_RESPONSE_TIMEOUT).
ClientResponseTimeout time.Duration ClientResponseTimeout time.Duration
@@ -61,6 +86,10 @@ type Config struct {
RateLimitPerMinute int64 RateLimitPerMinute int64
RateLimitPerHour int64 RateLimitPerHour int64
RateLimitPerDay int64 RateLimitPerDay int64
// RateLimitExemptPaths are the path prefixes whose requests the rate
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
// Each starts with /.
RateLimitExemptPaths []string
// DeniedCountries are the countries whose clients are refused // DeniedCountries are the countries whose clients are refused
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not // (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
// empty, are the only countries whose clients are let through // empty, are the only countries whose clients are let through
@@ -71,7 +100,8 @@ type Config struct {
// BanResponse is the status a refused client is answered with, 403 // BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer // or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that // (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
// breaks a rate limit, SWWAF_DENY_NETS and the country lists. // breaks a rate limit or matches a ban rule, SWWAF_DENY_NETS and the
// country lists.
BanResponse int BanResponse int
// LimitBanDuration is the ban for a first broken rate limit // LimitBanDuration is the ban for a first broken rate limit
// (SWWAF_LIMIT_BAN_DURATION). A limit broken again within // (SWWAF_LIMIT_BAN_DURATION). A limit broken again within
@@ -82,14 +112,78 @@ type Config struct {
LimitBanDuration time.Duration LimitBanDuration time.Duration
LimitBanRepeatWindow time.Duration LimitBanRepeatWindow time.Duration
MaxBanDuration time.Duration MaxBanDuration time.Duration
// AttackBanDuration is the ban for a first clear sign of attack
// (SWWAF_ATTACK_BAN_DURATION). It cannot be off.
AttackBanDuration time.Duration
// MaxBans is the most bans held (SWWAF_MAX_BANS). // MaxBans is the most bans held (SWWAF_MAX_BANS).
MaxBans int MaxBans int
// BanScopeV4Prefix is the length of the netblock around an IPv4 // BanScopeV4Prefix is the length of the netblock around an IPv4
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX). // client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
BanScopeV4Prefix int BanScopeV4Prefix int
// StateDir is the directory of the state files, an absolute path
// (SWWAF_STATE_DIR). bans.json is written StateWriteDelay after a ban
// is made (SWWAF_STATE_WRITE_DELAY), and every state file every
// StateCounterInterval (SWWAF_STATE_COUNTER_INTERVAL). Neither can be
// off.
StateDir string
StateWriteDelay time.Duration
StateCounterInterval time.Duration
// LogRequestHeaders are the request headers whose values the request
// log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS).
LogRequestHeaders []string
// AdminToken is the bearer token an admin sends for the ban endpoints
// and /_smallwebwaf/clients/<ip> (SWWAF_ADMIN_TOKEN), "" while it is
// unset and they are off.
AdminToken string
// MetricsToken is the bearer token a scraper sends for the metrics
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
// MetricsTopN is how many countries get series of their own in the
// metrics (SWWAF_METRICS_TOP_N).
MetricsToken string
MetricsTopN int
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
// unless RulesEnabled is false (SWWAF_RULES_ENABLED).
RulesDir string
RulesEnabled bool
// LogRemoteURL is where every line on stdout is also sent
// (SWWAF_LOG_REMOTE_URL), nil while it is unset and nothing is sent.
// LogRemoteTLSCAs are the certificates a syslog+tls endpoint's
// certificate must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil while
// it is unset and the host's own are used. LogRemoteBuffer is the most
// lines held while they wait to be sent (SWWAF_LOG_REMOTE_BUFFER).
// LogRemoteFacility is the number of the syslog facility
// (SWWAF_LOG_REMOTE_FACILITY), and LogRemoteAppName the APP-NAME
// (SWWAF_LOG_REMOTE_APP_NAME, by default InstanceName), of the records
// the lines are sent in.
LogRemoteURL *url.URL
LogRemoteTLSCAs *x509.CertPool
LogRemoteBuffer int
LogRemoteFacility int
LogRemoteAppName string
// AlertWebhookURL is where each alert is posted as JSON
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. AlertWebhookHeaders
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
// AlertSlackWebhookURL is the Slack incoming webhook each alert is
// posted to as a message (SWWAF_ALERT_SLACK_WEBHOOK_URL), and
// AlertNtfyURL the ntfy topic each is published to
// (SWWAF_ALERT_NTFY_URL), each nil while it is unset; AlertNtfyToken,
// unless empty, is sent to ntfy with each (SWWAF_ALERT_NTFY_TOKEN).
// With none of the three URLs set, no alert is sent. AlertEvents are
// the events alerts are sent for (SWWAF_ALERT_EVENTS). A repeat of an
// alert within AlertCooldown is held back (SWWAF_ALERT_COOLDOWN), and
// so is an alert past AlertMaxPerHour in an hour, for the hour's
// summary (SWWAF_ALERT_MAX_PER_HOUR); 0 is off for both.
AlertWebhookURL *url.URL
AlertWebhookHeaders http.Header
AlertSlackWebhookURL *url.URL
AlertNtfyURL *url.URL
AlertNtfyToken string
AlertEvents []string
AlertCooldown time.Duration
AlertMaxPerHour int
// settings are the values read, as given or by default, for the // settings are the values read, as given or by default, and the
// log line at start. // files they were read from, for the log line at start.
settings []slog.Attr settings []slog.Attr
} }
@@ -103,6 +197,15 @@ const (
mebibyte = 1 << 20 mebibyte = 1 << 20
gibibyte = 1 << 30 gibibyte = 1 << 30
ipv4Bits = 32 ipv4Bits = 32
// minTokenLength is the fewest characters a token may have.
minTokenLength = 32
// masked is what the log shows for a token that is set, and in place of
// a secret in another setting.
masked = "********"
// defaultListenAddr and defaultUpstreamURL are the defaults of
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
defaultListenAddr = ":8080"
defaultUpstreamURL = "http://127.0.0.1:8081"
) )
var ( var (
@@ -123,7 +226,13 @@ var (
"such as http://127.0.0.1:8081") "such as http://127.0.0.1:8081")
errNotCountry = errors.New( errNotCountry = errors.New(
"is not a two-letter country code such as de or kp") "is not a two-letter country code such as de or kp")
errNotHeaderName = errors.New(
"is not a header name such as accept-language")
errHeaderTakenOut = errors.New(
"is taken out of every request by Go's HTTP server, so it can never " +
"be logged")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too") errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New( errNotDurationAboveZero = errors.New(
"is not a duration above zero, such as 1h or 7d") "is not a duration above zero, such as 1h or 7d")
errNotNumberAboveZero = errors.New( errNotNumberAboveZero = errors.New(
@@ -131,18 +240,53 @@ var (
errNotBanResponse = errors.New("is not 403, 429 or close") errNotBanResponse = errors.New("is not 403, 429 or close")
errNotV4Prefix = errors.New( errNotV4Prefix = errors.New(
"is not the length of an IPv4 netblock, from 0 to 32, such as 24") "is not the length of an IPv4 netblock, from 0 to 32, such as 24")
errNotAbsolutePath = errors.New(
"is not an absolute path, such as /var/lib/smallwebwaf")
errShortToken = errors.New("is shorter than 32 characters")
errNotMode = errors.New("is not enforce or observe")
errNotPathPrefix = errors.New(
"is not a path prefix starting with /, such as /assets/")
errNotBoolean = errors.New("is not true or false")
errNotLogRemoteURL = errors.New(
"is not syslog+udp, syslog+tcp or syslog+tls with a host and a port, " +
"and nothing more, such as syslog+tls://logs.example:6514")
errNoCertificate = errors.New("holds no PEM certificate")
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
errNotAppName = errors.New(
"is not 1 to 48 printable ASCII characters without a space, such as gitea")
errSetTwice = errors.New("set only one of them")
errNotWebhookURL = errors.New(
"is not an http or https URL without a user or a fragment, " +
"such as https://alerts.example/smallwebwaf")
errNotWebhookHeader = errors.New(
"is not a header name followed by : and the header's value, " +
"such as Authorization:Bearer <token>")
errControlCharacter = errors.New(
"holds a control character, such as the carriage return of a Windows line end")
errNotAlertEvent = errors.New(
"is not ban, permanent_ban, waf_block, anomaly, reputation_hit, " +
"source_failure or file_error")
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
) )
// FromEnvironment reads the settings with lookupEnv, normally // FromEnvironment reads the settings with lookupEnv, normally
// os.LookupEnv. A setting that is not set takes its default. A setting // os.LookupEnv. A setting may instead be given as a file: the variable
// that is set but invalid is an error that names it. // named by the setting's name with _FILE added names the file, which is
// read now (see lookup). A setting that is not set takes its default. A
// setting that is set but invalid is an error that names it.
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) { func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv} env := &environment{lookupEnv: lookupEnv}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg := &Config{ cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"), ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"), UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges), InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"), Observe: env.observe("SWWAF_MODE", "enforce"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
ClientRequestHeaderMaxBytes: env.headerSize(
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES", "32K"),
ClientIdleTimeout: env.duration("SWWAF_CLIENT_IDLE_TIMEOUT", "120s"),
ClientResponseTimeout: env.duration("SWWAF_CLIENT_RESPONSE_TIMEOUT", "30m"), ClientResponseTimeout: env.duration("SWWAF_CLIENT_RESPONSE_TIMEOUT", "30m"),
UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"), UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"),
UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"), UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"),
@@ -154,6 +298,7 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"), RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"),
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"), RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"), RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""), DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries( ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""), "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
@@ -161,10 +306,39 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"), LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"), LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"), MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
AttackBanDuration: env.durationNotOff("SWWAF_ATTACK_BAN_DURATION", "7d"),
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"), MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"), BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
"accept,accept-language,accept-encoding,content-type,origin,range"),
AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
AlertWebhookURL: env.webhookURL("SWWAF_ALERT_WEBHOOK_URL"),
AlertWebhookHeaders: env.webhookHeaders("SWWAF_ALERT_WEBHOOK_HEADERS"),
AlertSlackWebhookURL: env.webhookURL("SWWAF_ALERT_SLACK_WEBHOOK_URL"),
AlertNtfyURL: env.webhookURL("SWWAF_ALERT_NTFY_URL"),
AlertNtfyToken: env.secret("SWWAF_ALERT_NTFY_TOKEN"),
AlertEvents: env.alertEvents("SWWAF_ALERT_EVENTS",
strings.Join(alerts.Events(), ",")),
AlertCooldown: env.duration("SWWAF_ALERT_COOLDOWN", "15m"),
AlertMaxPerHour: env.numberOrOff("SWWAF_ALERT_MAX_PER_HOUR", "60"),
} }
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
cfg.InstanceName, cfg.LogRemoteURL != nil)
env.checkInstanceNameForNtfy(cfg.InstanceName, cfg.AlertNtfyURL != nil)
for _, country := range cfg.ExclusivelyAllowedCountries { for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) { if slices.Contains(cfg.DeniedCountries, country) {
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
@@ -181,6 +355,24 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
return cfg, nil return cfg, nil
} }
// ListenAddrAndUpstreamURL reads only SWWAF_LISTEN_ADDR and
// SWWAF_UPSTREAM_URL, either of which may be given as a file, as
// FromEnvironment does. The health check needs no other setting, so it
// reads no other, nor a file that another names.
func ListenAddrAndUpstreamURL(
lookupEnv func(string) (string, bool),
) (string, *url.URL, error) {
env := &environment{lookupEnv: lookupEnv}
listenAddr := env.address("SWWAF_LISTEN_ADDR", defaultListenAddr)
upstreamURL := env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL)
if env.err != nil {
return "", nil, env.err
}
return listenAddr, upstreamURL, nil
}
// privateRanges are the private address ranges, the default trusted // privateRanges are the private address ranges, the default trusted
// proxies. // proxies.
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16" const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
@@ -202,8 +394,8 @@ type environment struct {
// value returns a setting's value, or its default when it is not set, // value returns a setting's value, or its default when it is not set,
// and notes it for the log. // and notes it for the log.
func (e *environment) value(name, defaultValue string) string { func (e *environment) value(name, defaultValue string) string {
value, ok := e.lookupEnv(name) value, set := e.lookup(name)
if !ok { if !set {
value = defaultValue value = defaultValue
} }
@@ -212,6 +404,37 @@ func (e *environment) value(name, defaultValue string) string {
return value return value
} }
// lookup returns a setting's value and whether it is set: the value of the
// variable name, or the contents of the file that the variable name_FILE
// names, less one newline at their end. It notes that file's path for the
// log. Both variables set, or a file that cannot be read, is an error.
func (e *environment) lookup(name string) (string, bool) {
value, set := e.lookupEnv(name)
fileName := name + "_FILE"
path, inFile := e.lookupEnv(fileName)
if !inFile {
return value, set
}
if set {
e.check(name, fmt.Errorf("is set, and so is %s; %w", fileName, errSetTwice))
return value, set
}
e.settings = append(e.settings, slog.String(fileName, path))
contents, err := os.ReadFile(path) //nolint:gosec // a file the admin names
if err != nil {
e.check(fileName, fmt.Errorf("cannot be read: %w", err))
return "", false
}
return strings.TrimSuffix(string(contents), "\n"), true
}
// check keeps the first error, naming the setting it is about. // check keeps the first error, naming the setting it is about.
func (e *environment) check(name string, err error) { func (e *environment) check(name string, err error) {
if err != nil && e.err == nil { if err != nil && e.err == nil {
@@ -235,6 +458,27 @@ func (e *environment) appURL(name, defaultValue string) *url.URL {
return upstream return upstream
} }
// observe reads the setting that is the mode, enforce or observe, and
// reports whether it is observe.
func (e *environment) observe(name, defaultValue string) bool {
mode := e.value(name, defaultValue)
if mode != "enforce" && mode != "observe" {
e.check(name, fmt.Errorf("%q %w", mode, errNotMode))
}
return mode == "observe"
}
// boolean reads a setting that is true or false.
func (e *environment) boolean(name, defaultValue string) bool {
value := e.value(name, defaultValue)
if value != "true" && value != "false" {
e.check(name, fmt.Errorf("%q %w", value, errNotBoolean))
}
return value == "true"
}
// netblocks reads a setting that is a list of netblocks. // netblocks reads a setting that is a list of netblocks.
func (e *environment) netblocks(name, defaultValue string) []netip.Prefix { func (e *environment) netblocks(name, defaultValue string) []netip.Prefix {
netblocks, err := parseNetblocks(e.value(name, defaultValue)) netblocks, err := parseNetblocks(e.value(name, defaultValue))
@@ -259,6 +503,15 @@ func (e *environment) size(name, defaultValue string) int64 {
return size return size
} }
// headerSize reads the setting that is the largest request line and
// headers.
func (e *environment) headerSize(name, defaultValue string) int64 {
size, err := parseHeaderSize(e.value(name, defaultValue))
e.check(name, err)
return size
}
// count reads a setting that is a number of requests. // count reads a setting that is a number of requests.
func (e *environment) count(name, defaultValue string) int64 { func (e *environment) count(name, defaultValue string) int64 {
count, err := parseCount(e.value(name, defaultValue)) count, err := parseCount(e.value(name, defaultValue))
@@ -267,6 +520,14 @@ func (e *environment) count(name, defaultValue string) int64 {
return count return count
} }
// pathPrefixes reads a setting that is a list of path prefixes.
func (e *environment) pathPrefixes(name, defaultValue string) []string {
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
e.check(name, err)
return prefixes
}
// countries reads a setting that is a list of countries. // countries reads a setting that is a list of countries.
func (e *environment) countries(name, defaultValue string) []string { func (e *environment) countries(name, defaultValue string) []string {
countries, err := parseCountries(e.value(name, defaultValue)) countries, err := parseCountries(e.value(name, defaultValue))
@@ -275,10 +536,19 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries return countries
} }
// headerNames reads a setting that is a list of header names, and
// returns them in lower case.
func (e *environment) headerNames(name, defaultValue string) []string {
headers, err := parseHeaderNames(e.value(name, defaultValue))
e.check(name, err)
return headers
}
// durationNotOff reads a setting that is a duration and, unlike a // durationNotOff reads a setting that is a duration and, unlike a
// timeout, cannot be off. // timeout, cannot be off.
func (e *environment) durationNotOff(name, defaultValue string) time.Duration { func (e *environment) durationNotOff(name, defaultValue string) time.Duration {
duration, err := parseDurationNotOff(e.value(name, defaultValue)) duration, err := ParseDurationNotOff(e.value(name, defaultValue))
e.check(name, err) e.check(name, err)
return duration return duration
@@ -309,6 +579,187 @@ func (e *environment) v4Prefix(name, defaultValue string) int {
return length return length
} }
// absolutePath reads a setting that is an absolute path.
func (e *environment) absolutePath(name, defaultValue string) string {
path := e.value(name, defaultValue)
if !filepath.IsAbs(path) {
e.check(name, fmt.Errorf("%q %w", path, errNotAbsolutePath))
}
return path
}
// token reads a setting that is a bearer token. Unset, it is "", which
// switches off what it guards; set, it must be at least minTokenLength
// characters. Neither the log nor an error shows its value.
func (e *environment) token(name string) string {
value, set := e.lookup(name)
if !set {
e.settings = append(e.settings, slog.String(name, ""))
return ""
}
e.settings = append(e.settings, slog.String(name, masked))
if utf8.RuneCountInString(value) < minTokenLength {
e.check(name, errShortToken)
}
return value
}
// logRemoteURL reads the setting that is where every log line is also
// sent. Unset or empty, it is nil, and nothing is sent.
func (e *environment) logRemoteURL(name string) *url.URL {
value := e.value(name, "")
if value == "" {
return nil
}
remote, err := parseLogRemoteURL(value)
e.check(name, err)
return remote
}
// certificates reads a setting that is the path of a file of PEM
// certificates. Unset or empty, it is nil. Its value names a file
// already, so, unlike the other settings, it has no _FILE form.
func (e *environment) certificates(name string) *x509.CertPool {
path, _ := e.lookupEnv(name)
e.settings = append(e.settings, slog.String(name, path))
if path == "" {
return nil
}
pem, err := os.ReadFile(path) //nolint:gosec // a file the admin names
if err != nil {
e.check(name, fmt.Errorf("cannot be read: %w", err))
return nil
}
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(pem) {
e.check(name, fmt.Errorf("%q %w", path, errNoCertificate))
return nil
}
return pool
}
// facility reads a setting that is a syslog facility, and returns its
// number.
func (e *environment) facility(name, defaultValue string) int {
number, err := parseFacility(e.value(name, defaultValue))
e.check(name, err)
return number
}
// appName reads the setting that is the APP-NAME of the records the log
// lines are sent in, by default the instance name. Its value is checked
// when it is set, and, while lines are sent, when it is the instance name.
func (e *environment) appName(name, instanceName string, sending bool) string {
value, set := e.lookup(name)
if !set {
value = instanceName
}
e.settings = append(e.settings, slog.String(name, value))
switch {
case isAppName(value):
case set:
e.check(name, fmt.Errorf("%q %w", value, errNotAppName))
case sending:
e.check(name, fmt.Errorf("is unset, and SWWAF_INSTANCE_NAME %q, its default, %w",
value, errNotAppName))
}
return value
}
// checkInstanceNameForNtfy refuses an instance name that holds a control
// character while ntfySet, SWWAF_ALERT_NTFY_URL being set: ntfy is sent
// the instance name in a header, which cannot hold one.
func (e *environment) checkInstanceNameForNtfy(instanceName string, ntfySet bool) {
if ntfySet && strings.ContainsFunc(instanceName, unicode.IsControl) {
e.check("SWWAF_INSTANCE_NAME", fmt.Errorf(
"%q %w, and is sent to ntfy in a header while SWWAF_ALERT_NTFY_URL is set",
instanceName, errControlCharacter))
}
}
// webhookURL reads a setting that is a URL each alert is posted to:
// SWWAF_ALERT_WEBHOOK_URL, SWWAF_ALERT_SLACK_WEBHOOK_URL or
// SWWAF_ALERT_NTFY_URL. Unset or empty, it is nil, and no alert is posted
// there. The log shows ******** in place of its path and query, and an
// error shows none of it, since a webhook or an ntfy topic can carry its
// secret there.
func (e *environment) webhookURL(name string) *url.URL {
value, _ := e.lookup(name)
webhook, logged, err := parseWebhookURL(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return webhook
}
// webhookHeaders reads the setting that is the headers sent with each
// alert. The log shows each header's value as ********, since a header
// such as Authorization carries a secret.
func (e *environment) webhookHeaders(name string) http.Header {
value, _ := e.lookup(name)
headers, logged, err := parseWebhookHeaders(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return headers
}
// secret reads a setting that is a secret another service gave, such as
// an ntfy token, "" while it is unset. It is sent in a header, which
// cannot hold a control character, so one in it is an error. The log
// shows ******** in place of a value that is not empty, and an error
// shows none of it.
func (e *environment) secret(name string) string {
value, _ := e.lookup(name)
logged := ""
if value != "" {
logged = masked
}
e.settings = append(e.settings, slog.String(name, logged))
if strings.ContainsFunc(value, unicode.IsControl) {
e.check(name, errControlCharacter)
}
return value
}
// alertEvents reads the setting that is the events alerts are sent for.
func (e *environment) alertEvents(name, defaultValue string) []string {
events, err := parseAlertEvents(e.value(name, defaultValue))
e.check(name, err)
return events
}
// numberOrOff reads a setting that is a whole number above zero, or off,
// which is 0.
func (e *environment) numberOrOff(name, defaultValue string) int {
number, err := parseNumberOrOff(e.value(name, defaultValue))
e.check(name, err)
return number
}
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a // parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
// whole number of days such as 7d, or off. // whole number of days such as 7d, or off.
func parseDuration(value string) (time.Duration, error) { func parseDuration(value string) (time.Duration, error) {
@@ -364,6 +815,19 @@ func parseSize(value string) (int64, error) {
return n * unit, nil return n * unit, nil
} }
// parseHeaderSize reads the largest request line and headers: a size as
// parseSize reads it, but more than 4K and never off. Go's server reads 4K
// past the limit it is given before it refuses, so proxy.New gives it this
// size less 4K, which must leave a limit.
func parseHeaderSize(value string) (int64, error) {
size, err := parseSize(value)
if err != nil || size <= 4*kibibyte {
return 0, fmt.Errorf("%q %w", value, errNotOver4K)
}
return size, nil
}
// splitUnit splits a size into its number and the bytes its suffix // splitUnit splits a size into its number and the bytes its suffix
// stands for. // stands for.
func splitUnit(value string) (string, int64) { func splitUnit(value string) (string, int64) {
@@ -397,9 +861,9 @@ func parseCount(value string) (int64, error) {
return n, nil return n, nil
} }
// parseDurationNotOff reads a duration above zero, as parseDuration does, // ParseDurationNotOff reads a duration above zero, as parseDuration does,
// but not off. // but not off. The ban endpoint reads the duration of a ban with it too.
func parseDurationNotOff(value string) (time.Duration, error) { func ParseDurationNotOff(value string) (time.Duration, error) {
duration, err := parseDuration(value) duration, err := parseDuration(value)
if err != nil || duration == 0 { if err != nil || duration == 0 {
return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero) return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero)
@@ -502,6 +966,23 @@ func parseNetblock(value string) (netip.Prefix, error) {
return netip.PrefixFrom(addr, addr.BitLen()), nil return netip.PrefixFrom(addr, addr.BitLen()), nil
} }
// parsePathPrefixes reads a comma-separated list of path prefixes, each
// starting with /.
func parsePathPrefixes(value string) ([]string, error) {
prefixes, err := parseList(value)
if err != nil {
return nil, err
}
for _, prefix := range prefixes {
if !strings.HasPrefix(prefix, "/") {
return nil, fmt.Errorf("%q %w", prefix, errNotPathPrefix)
}
}
return prefixes, nil
}
// countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK, // countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK,
// the code in common use for Kosovo. golang.org/x/text/language cannot // the code in common use for Kosovo. golang.org/x/text/language cannot
// check them: it also takes withdrawn codes such as su, and reserved ones // check them: it also takes withdrawn codes such as su, and reserved ones
@@ -558,6 +1039,58 @@ func parseCountries(value string) ([]string, error) {
return countries, nil return countries, nil
} }
// headerNameChars are the characters RFC 9110 allows in a header name:
// letters, digits and these marks.
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
"0123456789!#$%&'*+-.^_`|~"
// IsHeaderName reports whether name can be a header name: one or more of
// the characters RFC 9110 allows in one.
func IsHeaderName(name string) bool {
if name == "" {
return false
}
for _, char := range name {
if !strings.ContainsRune(headerNameChars, char) {
return false
}
}
return true
}
// parseHeaderNames reads a comma-separated list of header names in either
// case, and returns them in lower case. Host and Transfer-Encoding are
// refused: Go's HTTP server takes them out of the request's headers.
func parseHeaderNames(value string) ([]string, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
headers := make([]string, 0, len(items))
for _, item := range items {
if !IsHeaderName(item) {
return nil, fmt.Errorf("%q %w", item, errNotHeaderName)
}
header := strings.ToLower(item)
switch header {
case "host":
return nil, fmt.Errorf("%q %w; the request's host is the field host",
item, errHeaderTakenOut)
case "transfer-encoding":
return nil, fmt.Errorf("%q %w", item, errHeaderTakenOut)
}
headers = append(headers, header)
}
return headers, nil
}
// parseListenAddr checks an address to listen on: an optional host and a // parseListenAddr checks an address to listen on: an optional host and a
// port number. // port number.
func parseListenAddr(value string) (string, error) { func parseListenAddr(value string) (string, error) {
@@ -600,3 +1133,154 @@ func parseUpstreamURL(value string) (*url.URL, error) {
return upstream, nil return upstream, nil
} }
// parseLogRemoteURL reads where every log line is also sent:
// syslog+udp, syslog+tcp or syslog+tls, a host and a port from 1 to
// 65535, and nothing else.
func parseLogRemoteURL(value string) (*url.URL, error) {
remote, err := url.Parse(value)
if err != nil {
return nil, fmt.Errorf("%q %w", value, errNotLogRemoteURL)
}
schemes := []string{remotelog.SchemeUDP, remotelog.SchemeTCP, remotelog.SchemeTLS}
port, err := strconv.ParseUint(remote.Port(), 10, 16)
onlySchemeHostAndPort := slices.Contains(schemes, remote.Scheme) &&
remote.Hostname() != "" && err == nil && port != 0 &&
remote.User == nil && remote.Opaque == "" &&
(remote.Path == "" || remote.Path == "/") &&
remote.RawQuery == "" && remote.Fragment == ""
if !onlySchemeHostAndPort {
return nil, fmt.Errorf("%q %w", value, errNotLogRemoteURL)
}
return remote, nil
}
// parseFacility reads the name of a syslog facility, and returns its
// number, as RFC 5424 numbers them.
func parseFacility(value string) (int, error) {
//nolint:mnd // the facilities' numbers in RFC 5424
number, known := map[string]int{
"kern": 0, "user": 1, "mail": 2, "daemon": 3, "auth": 4, "syslog": 5,
"lpr": 6, "news": 7, "uucp": 8, "cron": 9, "authpriv": 10, "ftp": 11,
"local0": 16, "local1": 17, "local2": 18, "local3": 19,
"local4": 20, "local5": 21, "local6": 22, "local7": 23,
}[value]
if !known {
return 0, fmt.Errorf("%q %w", value, errNotFacility)
}
return number, nil
}
// parseWebhookURL reads where each alert is posted: http or https, a
// host, and an optional port from 1 to 65535, path and query, without a
// user or a fragment. It returns the URL, and how the log shows it: its
// scheme and host, and ******** in place of its path and query, if it has
// either. An error shows no part of the value. An empty value is no URL.
func parseWebhookURL(value string) (*url.URL, string, error) {
if value == "" {
return nil, "", nil
}
webhook, err := url.Parse(value)
if err != nil {
return nil, "", errNotWebhookURL
}
port, err := strconv.ParseUint(webhook.Port(), 10, 16)
valid := (webhook.Scheme == "http" || webhook.Scheme == "https") &&
webhook.Hostname() != "" && (webhook.Port() == "" || (err == nil && port != 0)) &&
webhook.User == nil && webhook.Opaque == "" && webhook.Fragment == ""
if !valid {
return nil, "", errNotWebhookURL
}
logged := webhook.Scheme + "://" + webhook.Host
if webhook.Path != "" || webhook.RawQuery != "" {
logged += "/" + masked
}
return webhook, logged, nil
}
// parseWebhookHeaders reads a comma-separated list of headers, each its
// name, :, and its value, and returns them, and how the log shows them,
// with each value as ********. An error names the item by its place in
// the list, so that it shows no value. An empty value is an empty list.
func parseWebhookHeaders(value string) (http.Header, string, error) {
headers := http.Header{}
if strings.TrimSpace(value) == "" {
return headers, "", nil
}
logged := []string{}
for i, item := range strings.Split(value, ",") {
name, headerValue, found := strings.Cut(item, ":")
name = strings.TrimSpace(name)
if !found || !IsHeaderName(name) || strings.ContainsAny(headerValue, "\r\n\x00") {
return nil, "", fmt.Errorf("item %d %w", i+1, errNotWebhookHeader)
}
headers.Add(name, strings.TrimSpace(headerValue))
logged = append(logged, name+":"+masked)
}
return headers, strings.Join(logged, ","), nil
}
// parseAlertEvents reads a comma-separated list of the events alerts can
// be sent for.
func parseAlertEvents(value string) ([]string, error) {
events, err := parseList(value)
if err != nil {
return nil, err
}
for _, event := range events {
if !slices.Contains(alerts.Events(), event) {
return nil, fmt.Errorf("%q %w", event, errNotAlertEvent)
}
}
return events, nil
}
// parseNumberOrOff reads a whole number above zero, or off, which is 0.
func parseNumberOrOff(value string) (int, error) {
if value == off {
return 0, nil
}
n, err := strconv.Atoi(value)
if err != nil || n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotNumberOrOff)
}
return n, nil
}
// appNameMaxLength is the most characters RFC 5424 allows in an
// APP-NAME.
const appNameMaxLength = 48
// isAppName reports whether value can be an APP-NAME: 1 to
// appNameMaxLength printable ASCII characters, none of them a space.
func isAppName(value string) bool {
if value == "" || len(value) > appNameMaxLength {
return false
}
for _, char := range []byte(value) {
if char < '!' || char > '~' {
return false
}
}
return true
}
File diff suppressed because it is too large Load Diff
+9
View File
@@ -0,0 +1,9 @@
package lookup
import "net/http"
// SetTransport has g's requests to GeoJS go through transport instead of
// the network.
func (g *GeoJS) SetTransport(transport http.RoundTripper) {
g.httpClient.Transport = transport
}
+97 -13
View File
@@ -1,6 +1,7 @@
// Package lookup looks up each client's country through the GeoJS web // Package lookup looks up each client's country through the GeoJS web
// service, and keeps the answers in memory, for at most 100,000 clients // service, and keeps the answers in memory, for at most 100,000 clients
// and for 7 days each. // and for 7 days each. The answers are written to lookups.json and read
// from it by the state package.
package lookup package lookup
import ( import (
@@ -12,11 +13,14 @@ import (
"log/slog" "log/slog"
"net/http" "net/http"
"net/netip" "net/netip"
"slices"
"strings" "strings"
"sync" "sync"
"time" "time"
"github.com/hashicorp/golang-lru/v2/simplelru" "github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics"
) )
// URL is GeoJS's country endpoint. Asked about several addresses at once, // URL is GeoJS's country endpoint. Asked about several addresses at once,
@@ -62,6 +66,11 @@ type Params struct {
Now func() time.Time Now func() time.Time
// ProcessLog receives GeoJS's failures. // ProcessLog receives GeoJS's failures.
ProcessLog *slog.Logger ProcessLog *slog.Logger
// Metrics count the requests to GeoJS, those that failed, and the
// clients that go without an answer.
Metrics *metrics.Metrics
// Alerts receive a source_failure alert each time GeoJS fails.
Alerts *alerts.Queue
} }
// GeoJS looks up clients' countries through GeoJS. At most one request // GeoJS looks up clients' countries through GeoJS. At most one request
@@ -71,12 +80,14 @@ type GeoJS struct {
url string url string
now func() time.Time now func() time.Time
processLog *slog.Logger processLog *slog.Logger
metrics *metrics.Metrics
alerts *alerts.Queue
// httpClient follows no redirect, so that visitors' addresses go to // httpClient follows no redirect, so that visitors' addresses go to
// GeoJS alone: a redirect is a failure. // GeoJS alone: a redirect is a failure.
httpClient *http.Client httpClient *http.Client
mu sync.Mutex mu sync.Mutex
answers *simplelru.LRU[netip.Prefix, answer] answers *simplelru.LRU[netip.Prefix, *Answer]
// waiting are the clients without an answer: those to ask GeoJS about, // waiting are the clients without an answer: those to ask GeoJS about,
// and those it is being asked about. // and those it is being asked about.
waiting map[netip.Prefix]*wait waiting map[netip.Prefix]*wait
@@ -88,11 +99,14 @@ type GeoJS struct {
retryAt time.Time retryAt time.Time
} }
// answer is what GeoJS said about a client: its country, "" when GeoJS // Answer is what GeoJS said about a client, as lookups.json holds it: its
// cannot place it, and when GeoJS said so. // country, "" when GeoJS cannot place it, when GeoJS said so, and when
type answer struct { // the answer was last used.
country string type Answer struct {
received time.Time Client netip.Prefix `json:"client"`
Country string `json:"country"`
Answered time.Time `json:"answered"`
Used time.Time `json:"used"`
} }
// wait is a client waiting for its answer. // wait is a client waiting for its answer.
@@ -107,7 +121,7 @@ type wait struct {
// New returns a GeoJS with no answer kept yet. // New returns a GeoJS with no answer kept yet.
func New(params Params) *GeoJS { func New(params Params) *GeoJS {
answers, err := simplelru.NewLRU[netip.Prefix, answer](maxAnswers, nil) answers, err := simplelru.NewLRU[netip.Prefix, *Answer](maxAnswers, 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
} }
@@ -116,6 +130,8 @@ func New(params Params) *GeoJS {
url: params.URL, url: params.URL,
now: params.Now, now: params.Now,
processLog: params.ProcessLog, processLog: params.ProcessLog,
metrics: params.Metrics,
alerts: params.Alerts,
httpClient: &http.Client{ httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error { CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse return http.ErrUseLastResponse
@@ -155,6 +171,9 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
defer g.mu.Unlock() defer g.mu.Unlock()
country, found := g.kept(client) country, found := g.kept(client)
if !found {
g.metrics.GeoJSUnanswered.Inc()
}
w, waiting := g.waiting[client] w, waiting := g.waiting[client]
if !found && waiting { if !found && waiting {
@@ -164,6 +183,49 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
return country return country
} }
// Snapshot returns every answer kept, sorted by client, as lookups.json
// lists them.
func (g *GeoJS) Snapshot() []Answer {
g.mu.Lock()
answers := make([]Answer, 0, g.answers.Len())
for _, kept := range g.answers.Values() {
answers = append(answers, *kept)
}
g.mu.Unlock()
slices.SortFunc(answers, func(a, b Answer) int {
return a.Client.Compare(b.Client)
})
return answers
}
// Load keeps answers read from lookups.json, in place of the answers it
// keeps, in the order they were last used, so that the one used longest
// ago is dropped first. Answers GeoJS gave keepFor ago or more are
// dropped.
func (g *GeoJS) Load(answers []Answer) {
answers = slices.Clone(answers)
slices.SortStableFunc(answers, func(a, b Answer) int {
return a.Used.Compare(b.Used)
})
g.mu.Lock()
defer g.mu.Unlock()
g.answers.Purge()
now := g.now()
for _, answer := range answers {
if now.Sub(answer.Answered) < keepFor {
g.answers.Add(answer.Client, &answer)
}
}
}
// answerOrWait returns client's kept answer if it has one. Otherwise it // answerOrWait returns client's kept answer if it has one. Otherwise it
// puts the client among those waiting if there is room, has GeoJS asked // puts the client among those waiting if there is room, has GeoJS asked
// about them if it can be, and returns what to wait on for the answer, or // about them if it can be, and returns what to wait on for the answer, or
@@ -188,6 +250,8 @@ func (g *GeoJS) answerOrWait(
g.ask(ctx) g.ask(ctx)
if w == nil { if w == nil {
g.metrics.GeoJSUnanswered.Inc()
return "", nil // too many clients wait already return "", nil // too many clients wait already
} }
@@ -197,21 +261,27 @@ func (g *GeoJS) answerOrWait(
} }
if w.late { if w.late {
g.metrics.GeoJSUnanswered.Inc()
return "", nil return "", nil
} }
return "", w.asked return "", w.asked
} }
// kept returns client's answer, if one was received less than keepFor // kept returns client's answer, if GeoJS gave it less than keepFor ago,
// ago. // and notes that it was used.
func (g *GeoJS) kept(client netip.Prefix) (string, bool) { func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
now := g.now()
kept, found := g.answers.Get(client) kept, found := g.answers.Get(client)
if !found || g.now().Sub(kept.received) >= keepFor { if !found || now.Sub(kept.Answered) >= keepFor {
return "", false return "", false
} }
return kept.country, true kept.Used = now
return kept.Country, true
} }
// ask starts asking GeoJS about the waiting clients, unless a request to // ask starts asking GeoJS about the waiting clients, unless a request to
@@ -293,7 +363,9 @@ func (g *GeoJS) keep(
continue continue
} }
g.answers.Add(client, answer{country: country, received: now}) g.answers.Add(client, &Answer{
Client: client, Country: country, Answered: now, Used: now,
})
close(g.waiting[client].asked) close(g.waiting[client].asked)
delete(g.waiting, client) delete(g.waiting, client)
} }
@@ -303,6 +375,8 @@ func (g *GeoJS) keep(
} }
if err != nil { if err != nil {
g.metrics.GeoJSFailures.Inc()
g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay), g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay),
maxRetryDelay) maxRetryDelay)
g.retryAt = now.Add(g.retryDelay) g.retryAt = now.Add(g.retryDelay)
@@ -317,6 +391,14 @@ func (g *GeoJS) keep(
g.processLog.Warn("asking GeoJS failed", g.processLog.Warn("asking GeoJS failed",
"error", err.Error(), "asking_again_in", g.retryDelay.String()) "error", err.Error(), "asking_again_in", g.retryDelay.String())
g.alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs", "error": err.Error(),
"asking_again_in": g.retryDelay.String(),
},
})
return false return false
} }
@@ -347,6 +429,8 @@ func (g *GeoJS) request(
req.URL.RawQuery = "ip=" + strings.Join(addrs, ",") req.URL.RawQuery = "ip=" + strings.Join(addrs, ",")
g.metrics.GeoJSRequests.Inc()
res, err := g.httpClient.Do(req) res, err := g.httpClient.Do(req)
if err != nil { if err != nil {
// Do's error names the URL, and so the visitors' addresses, which // Do's error names the URL, and so the visitors' addresses, which
+367 -229
View File
@@ -6,13 +6,19 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
"net/url"
"reflect"
"slices" "slices"
"strings" "strings"
"sync" "sync"
"testing" "testing"
"testing/synctest"
"time" "time"
"github.com/prometheus/client_golang/prometheus/testutil"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
) )
const ( const (
@@ -26,80 +32,87 @@ const (
leftOut = "203.0.113.7" leftOut = "203.0.113.7"
// timeout is how long a new client waits for its answer. // timeout is how long a new client waits for its answer.
timeout = time.Second timeout = time.Second
// waitLimit bounds how long a test waits for what should happen.
waitLimit = 10 * time.Second
// pollInterval is how often a test looks again.
pollInterval = 10 * time.Millisecond
// week is how long an answer is kept. // week is how long an answer is kept.
week = 7 * 24 * time.Hour week = 7 * 24 * time.Hour
) )
// The tests that have GeoJS asked run in a synctest bubble, where the time
// package runs on a clock of the test's own: a wait lasts exactly as long
// as it should, however slowly the test process runs, and synctest.Wait
// returns once g has done all it can before time passes. The stand-in for
// GeoJS answers without the network, since a request waiting on the
// network would keep that clock from moving on.
func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) { func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
placed := netip.MustParsePrefix("203.0.113.9/32") geojs, clock, g := start()
notPlaced := netip.MustParsePrefix(unplaced + "/32") placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "") wantCountry(t, g, notPlaced, "")
wantRequests(t, geojs, 2) wantRequests(t, geojs, 2)
// An answer without a country is kept too. // An answer without a country is kept too.
clock.advance(week - time.Second) clock.advance(week - time.Second)
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "") wantCountry(t, g, notPlaced, "")
wantRequests(t, geojs, 2) wantRequests(t, geojs, 2)
clock.advance(time.Second) clock.advance(time.Second)
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
wantRequests(t, geojs, 3) wantRequests(t, geojs, 3)
wantAsked(t, geojs, 2, "203.0.113.9") wantAsked(t, geojs, 2, "203.0.113.9")
})
} }
func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) { func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32") geojs, clock, g := start()
client := netip.MustParsePrefix("203.0.113.9/32")
// The client comes while GeoJS is asked about an earlier client, which // The client comes while GeoJS is asked about an earlier client, which
// it answers most of a second later. It is then asked about the client // it answers most of a second later. It is then asked about the client
// and does not answer: that request is abandoned a second after it // and does not answer: that request is abandoned a second after it
// began, well after the client's wait is over. // began, well after the client's wait is over.
geojs.set(answeringSlowly) geojs.set(answeringSlowly)
var earlier sync.WaitGroup var earlier sync.WaitGroup
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) }) earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
defer earlier.Wait() defer earlier.Wait()
waitForRequests(t, geojs, 1) waitForRequests(t, geojs, 1)
geojs.set(hanging) geojs.set(hanging)
began := time.Now() began := time.Now()
wantCountry(t, g, client, "") wantCountry(t, g, client, "")
took := time.Since(began) took := time.Since(began)
if took < timeout || took > timeout+timeout/2 { if took != timeout {
t.Errorf("waited %s for the answer, want %s", took, timeout) t.Errorf("waited %s for the answer, want %s", took, timeout)
} }
// Its next request does not wait. // Its next request does not wait.
began = time.Now() began = time.Now()
wantCountry(t, g, client, "") wantCountry(t, g, client, "")
took = time.Since(began) took = time.Since(began)
if took > timeout/2 { if took != 0 {
t.Errorf("waited %s again, want no wait", took) t.Errorf("waited %s again, want no wait", took)
} }
// Once GeoJS answers, the client is asked about again in the // Once GeoJS answers, the client is asked about again in the
// background, and has its country. // background, and has its country.
geojs.set(answering) geojs.set(answering)
waitForCountry(t, g, clock, client, germany) waitForCountry(t, g, clock, client, germany)
})
} }
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) { func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
@@ -118,33 +131,35 @@ func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
other := netip.MustParsePrefix("203.0.113.1/32") geojs, clock, g := start()
client := netip.MustParsePrefix(leftOut + "/32") other := netip.MustParsePrefix("203.0.113.1/32")
client := netip.MustParsePrefix(leftOut + "/32")
// GeoJS fails, and is left alone for a second while the client // GeoJS fails, and is left alone for a second while the client
// comes too, so that the next request asks about both. // comes too, so that the next request asks about both.
geojs.set(failing) geojs.set(failing)
wantCountry(t, g, other, "") wantCountry(t, g, other, "")
wantCountry(t, g, client, "") wantCountry(t, g, client, "")
geojs.set(tc.answers) geojs.set(tc.answers)
clock.advance(time.Second) clock.advance(time.Second)
wantCountry(t, g, other, "") wantCountry(t, g, other, "")
waitForRequests(t, geojs, 2) waitForRequests(t, geojs, 2)
// The answer counts as a failure, and the client is asked about // The answer counts as a failure, and the client is asked about
// again, with the other client only if the answer left it out too. // again, with the other client only if the answer left it out too.
geojs.set(answering) geojs.set(answering)
waitForCountry(t, g, clock, client, germany) waitForCountry(t, g, clock, client, germany)
wantCountry(t, g, other, germany) wantCountry(t, g, other, germany)
wantRequests(t, geojs, 3) wantRequests(t, geojs, 3)
if tc.named { if tc.named {
wantAsked(t, geojs, 2, leftOut) wantAsked(t, geojs, 2, leftOut)
} else { } else {
wantAsked(t, geojs, 2, leftOut, "203.0.113.1") wantAsked(t, geojs, 2, leftOut, "203.0.113.1")
} }
})
}) })
} }
} }
@@ -152,173 +167,266 @@ func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
func TestRedirectCountsAsFailure(t *testing.T) { func TestRedirectCountsAsFailure(t *testing.T) {
t.Parallel() t.Parallel()
geojs, _, g := start(t) synctest.Test(t, func(t *testing.T) {
geojs.set(redirecting) geojs, _, g := start()
geojs.set(redirecting)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
wantRequests(t, geojs, 1) wantRequests(t, geojs, 1)
})
} }
func TestCountryIsKeptInCapitals(t *testing.T) { func TestCountryIsKeptInCapitals(t *testing.T) {
t.Parallel() t.Parallel()
geojs, _, g := start(t) synctest.Test(t, func(t *testing.T) {
geojs.set(answeringInLowerCase) geojs, _, g := start()
geojs.set(answeringInLowerCase)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany) wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany)
})
} }
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) { func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
t.Parallel() t.Parallel()
var log strings.Builder synctest.Test(t, func(t *testing.T) {
var log strings.Builder
// Nothing listens on port 1, so asking GeoJS fails. // GeoJS does not answer, so the request to it is abandoned, and fails.
g := lookup.New(lookup.Params{ geojs := &standIn{answers: hanging}
URL: "http://127.0.0.1:1", g := lookup.New(lookup.Params{
Now: time.Now, URL: lookup.URL,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1),
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(geojs)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
synctest.Wait()
logged := log.String()
if !strings.Contains(logged, "asking GeoJS failed") ||
strings.Contains(logged, "203.0.113.9") {
t.Errorf("logged %q, want the failure without the address asked about", logged)
}
}) })
}
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") func TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
t.Parallel()
logged := log.String() synctest.Test(t, func(t *testing.T) {
if !strings.Contains(logged, "asking GeoJS failed") || geojs, clock, g, queue := startWithAlerts()
strings.Contains(logged, "203.0.113.9") { clients := newClients()
t.Errorf("logged %q, want the failure without the address asked about", logged)
} geojs.set(failing)
wantCountry(t, g, clients(), "")
want := alerts.Alert{
Time: clock.Now(),
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs",
"error": "GeoJS answered 503 Service Unavailable",
"asking_again_in": "1s",
},
}
// The next failure, a second later, is a repeat within the
// cooldown.
clock.advance(time.Second)
wantCountry(t, g, clients(), "")
wantRequests(t, geojs, 2)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
}
if queue.Suppressed() != 1 {
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
}
})
} }
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) { func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
// GeoJS fails, and is then left alone for a second, while three more // GeoJS fails, and is then left alone for a second, while three more
// clients come. An IPv6 client is a /64, and GeoJS is asked about its // clients come. An IPv6 client is a /64, and GeoJS is asked about its
// first address. // first address.
geojs.set(failing) geojs.set(failing)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "") wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "")
wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "") wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "")
wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "") wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "")
wantRequests(t, geojs, 1) wantRequests(t, geojs, 1)
geojs.set(answering) geojs.set(answering)
clock.advance(time.Second) clock.advance(time.Second)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany) wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany)
wantRequests(t, geojs, 2) wantRequests(t, geojs, 2)
wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3") wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3")
})
} }
func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) { func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
clients := newClients() geojs, clock, g := start()
kept := clients() clients := newClients()
kept := clients()
wantCountry(t, g, kept, germany)
geojs.set(failing)
wantCountry(t, g, kept, germany)
wantRequests(t, geojs, 1)
// Each failure leaves GeoJS alone twice as long as the one before, up
// to five minutes. New clients meanwhile count as not found, and the
// client with a kept answer still gets its country, without GeoJS being
// asked.
requests := 1
for _, delay := range []time.Duration{
time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second,
16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second,
256 * time.Second, 5 * time.Minute, 5 * time.Minute,
} {
wantCountry(t, g, clients(), "")
requests++
wantRequests(t, geojs, requests)
clock.advance(delay - time.Millisecond)
wantCountry(t, g, clients(), "")
wantCountry(t, g, kept, germany) wantCountry(t, g, kept, germany)
wantRequests(t, geojs, requests)
clock.advance(time.Millisecond) geojs.set(failing)
} wantCountry(t, g, kept, germany)
wantRequests(t, geojs, 1)
// Once GeoJS answers again, it is asked about every client waiting. // Each failure leaves GeoJS alone twice as long as the one before, up
geojs.set(answering) // to five minutes. New clients meanwhile count as not found, and the
wantCountry(t, g, clients(), germany) // client with a kept answer still gets its country, without GeoJS being
wantRequests(t, geojs, requests+1) // asked.
requests := 1
asked := waitForRequests(t, geojs, requests+1) for _, delay := range []time.Duration{
if len(asked[requests]) != 23 { time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second,
t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests])) 16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second,
} 256 * time.Second, 5 * time.Minute, 5 * time.Minute,
} {
wantCountry(t, g, clients(), "")
requests++
wantRequests(t, geojs, requests)
clock.advance(delay - time.Millisecond)
wantCountry(t, g, clients(), "")
wantCountry(t, g, kept, germany)
wantRequests(t, geojs, requests)
clock.advance(time.Millisecond)
}
// Once GeoJS answers again, it is asked about every client waiting.
geojs.set(answering)
wantCountry(t, g, clients(), germany)
wantRequests(t, geojs, requests+1)
asked := waitForRequests(t, geojs, requests+1)
if len(asked[requests]) != 23 {
t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests]))
}
})
} }
func TestAtMost200AddressesInOneRequest(t *testing.T) { func TestAtMost200AddressesInOneRequest(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
clients := newClients() geojs, clock, g := start()
first := clients() clients := newClients()
first := clients()
// 201 clients wait while GeoJS is left alone after a failure. // 201 clients wait while GeoJS is left alone after a failure.
geojs.set(failing) geojs.set(failing)
wantCountry(t, g, first, "") wantCountry(t, g, first, "")
for range 200 { for range 200 {
wantCountry(t, g, clients(), "") wantCountry(t, g, clients(), "")
} }
// The first one's next request has GeoJS asked again. // The first one's next request has GeoJS asked again.
geojs.set(answering) geojs.set(answering)
clock.advance(time.Second) clock.advance(time.Second)
wantCountry(t, g, first, "") wantCountry(t, g, first, "")
asked := waitForRequests(t, geojs, 3) asked := waitForRequests(t, geojs, 3)
if len(asked[1]) != 200 || len(asked[2]) != 1 { if len(asked[1]) != 200 || len(asked[2]) != 1 {
t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1", t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1",
len(asked[1]), len(asked[2])) len(asked[1]), len(asked[2]))
} }
})
} }
func TestAtMost10000ClientsWait(t *testing.T) { func TestAtMost10000ClientsWait(t *testing.T) {
t.Parallel() t.Parallel()
geojs, clock, g := start(t) synctest.Test(t, func(t *testing.T) {
clients := newClients() geojs, clock, g := start()
first := clients() clients := newClients()
first := clients()
// 10,000 clients wait while GeoJS is left alone after a failure, and // 10,000 clients wait while GeoJS is left alone after a failure, and
// one more cannot join them. // one more cannot join them.
geojs.set(failing) geojs.set(failing)
wantCountry(t, g, first, "") wantCountry(t, g, first, "")
for range 9999 { for range 9999 {
wantCountry(t, g, clients(), "") wantCountry(t, g, clients(), "")
}
extra := clients()
wantCountry(t, g, extra, "")
// The first one's next request has GeoJS asked about the 10,000, 200
// at a time, and not about the one more.
geojs.set(answering)
clock.advance(time.Second)
wantCountry(t, g, first, "")
asked := waitForRequests(t, geojs, 51)
for i, request := range asked {
if slices.Contains(request, extra.Addr().String()) {
t.Errorf("request %d asked about %s", i, extra.Addr())
} }
}
// With room among those waiting, it is asked about. extra := clients()
wantCountry(t, g, extra, germany) wantCountry(t, g, extra, "")
// The first one's next request has GeoJS asked about the 10,000, 200
// at a time, and not about the one more.
geojs.set(answering)
clock.advance(time.Second)
wantCountry(t, g, first, "")
asked := waitForRequests(t, geojs, 51)
for i, request := range asked {
if slices.Contains(request, extra.Addr().String()) {
t.Errorf("request %d asked about %s", i, extra.Addr())
}
}
// With room among those waiting, it is asked about.
wantCountry(t, g, extra, germany)
})
}
func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
m := metrics.New(1)
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(&standIn{answers: failing})
clients := newClients()
// GeoJS fails, so the first client goes without an answer, and GeoJS
// is left alone for a second, which does not pass in this test.
wantCountry(t, g, clients(), "")
wantUnanswered(t, m, 1)
// Meanwhile each new client goes without one at once, while there is
// room for it among the 10,000 that may wait.
for range 9999 {
wantCountry(t, g, clients(), "")
}
wantUnanswered(t, m, 10000)
// One more, for which there is no room, goes without one too.
wantCountry(t, g, clients(), "")
wantUnanswered(t, m, 10001)
})
} }
// How the stand-in for GeoJS answers. // How the stand-in for GeoJS answers.
@@ -337,13 +445,25 @@ const (
// standIn is a stand-in for GeoJS. It notes the addresses each request // standIn is a stand-in for GeoJS. It notes the addresses each request
// asks about. // asks about.
type standIn struct { type standIn struct {
server *httptest.Server
mu sync.Mutex mu sync.Mutex
answers int answers int
requests [][]string requests [][]string
} }
// RoundTrip has the stand-in answer req, in place of the network. A request
// abandoned before the stand-in answers fails, as over the network.
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
answer := httptest.NewRecorder()
s.ServeHTTP(answer, req)
err := req.Context().Err()
if err != nil {
return nil, err
}
return answer.Result(), nil
}
// ServeHTTP answers a request about the addresses in its ip parameter. // ServeHTTP answers a request about the addresses in its ip parameter.
func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
addrs := strings.Split(r.URL.Query().Get("ip"), ",") addrs := strings.Split(r.URL.Query().Get("ip"), ",")
@@ -422,7 +542,8 @@ func (s *standIn) asked() [][]string {
return slices.Clone(s.requests) return slices.Clone(s.requests)
} }
// testClock is a clock the test sets. // testClock is a clock the test sets. GeoJS tells the time by it, while
// waits run on the bubble's clock.
type testClock struct { type testClock struct {
mu sync.Mutex mu sync.Mutex
now time.Time now time.Time
@@ -444,25 +565,38 @@ func (c *testClock) advance(d time.Duration) {
c.now = c.now.Add(d) c.now = c.now.Add(d)
} }
// start starts a stand-in for GeoJS that answers, and returns it, a // start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
// clock, and a GeoJS asking it by that clock. // asking the stand-in by that clock.
func start(t *testing.T) (*standIn, *testClock, *lookup.GeoJS) { func start() (*standIn, *testClock, *lookup.GeoJS) {
t.Helper() geojs, clock, g, _ := startWithAlerts()
geojs := &standIn{}
geojs.server = httptest.NewServer(geojs)
t.Cleanup(geojs.server.Close)
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
g := lookup.New(lookup.Params{
URL: geojs.server.URL,
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
})
return geojs, clock, g return geojs, clock, g
} }
// startWithAlerts is start, and returns the queue of the alerts GeoJS
// raises as well, for a webhook that is never sent them, with the default
// cooldown, by the same clock.
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
geojs := &standIn{}
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: clock.Now,
})
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1),
Alerts: queue,
})
g.SetTransport(geojs)
return geojs, clock, g, queue
}
// newClients returns what returns a new IPv4 client each time it is // newClients returns what returns a new IPv4 client each time it is
// called. // called.
func newClients() func() netip.Prefix { func newClients() func() netip.Prefix {
@@ -513,41 +647,45 @@ func wantAsked(t *testing.T, geojs *standIn, i int, want ...string) {
} }
} }
// waitForRequests waits for GeoJS to have had count requests, and returns // wantUnanswered checks how many requests m counts as having gone without
// the addresses each asked about. // an answer from GeoJS.
func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
t.Helper()
got := testutil.ToFloat64(m.GeoJSUnanswered)
if got != want {
t.Errorf("%v requests went without an answer, want %v", got, want)
}
}
// waitForRequests waits until g has done all it can before time passes,
// checks that GeoJS has had count requests, and returns the addresses each
// asked about.
func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string { func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string {
t.Helper() t.Helper()
deadline := time.Now().Add(waitLimit) synctest.Wait()
for time.Now().Before(deadline) {
asked := geojs.asked()
if len(asked) >= count {
return asked
}
time.Sleep(pollInterval) asked := geojs.asked()
if len(asked) != count {
t.Fatalf("GeoJS had %d requests, want %d", len(asked), count)
} }
t.Fatalf("fewer than %d requests to GeoJS after %s", count, waitLimit) return asked
return nil
} }
// waitForCountry waits for g to give client the country want, moving the // waitForCountry lets a request to GeoJS under way be abandoned, and moves
// clock on a minute at a time, so that GeoJS is asked again after a // the clock on a minute, so that GeoJS may be asked again after a failure.
// failure. // It then checks that client's next request does not wait but has it asked
// about again in the background, after which g gives it the country want.
func waitForCountry( func waitForCountry(
t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string, t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string,
) { ) {
t.Helper() t.Helper()
deadline := time.Now().Add(waitLimit) time.Sleep(timeout)
for g.Country(t.Context(), client) != want { clock.advance(time.Minute)
if time.Now().After(deadline) { wantCountry(t, g, client, "")
t.Fatalf("%s is not in %q after %s", client, want, waitLimit) synctest.Wait()
} wantCountry(t, g, client, want)
clock.advance(time.Minute)
time.Sleep(pollInterval)
}
} }
+99
View File
@@ -0,0 +1,99 @@
package lookup_test
import (
"net/netip"
"slices"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/lookup"
)
func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
_, clock, g := start()
placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
asked := clock.Now()
wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "")
clock.advance(time.Hour)
wantCountry(t, g, placed, germany)
want := []lookup.Answer{
{Client: notPlaced, Country: "", Answered: asked, Used: asked},
{Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
}
if got := g.Snapshot(); !slices.Equal(got, want) {
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
}
})
}
func TestLoadedAnswersAreKeptFor7DaysFromWhenGeoJSGaveThem(t *testing.T) {
t.Parallel()
geojs, clock, g := start()
now := clock.Now()
kept := lookup.Answer{
Client: netip.MustParsePrefix("203.0.113.9/32"),
Country: "FR",
Answered: now.Add(-week + time.Second),
Used: now.Add(-time.Hour),
}
stale := lookup.Answer{
Client: netip.MustParsePrefix("203.0.113.10/32"),
Country: "FR",
Answered: now.Add(-week),
Used: now.Add(-time.Hour),
}
g.Load([]lookup.Answer{kept, stale})
if got := g.Snapshot(); !slices.Equal(got, []lookup.Answer{kept}) {
t.Errorf("kept %+v, want only the answer GeoJS gave less than 7 days ago", got)
}
wantCountry(t, g, kept.Client, "FR")
wantRequests(t, geojs, 0)
}
func TestLoadDropsTheAnswerUsedLongestAgoFirst(t *testing.T) {
t.Parallel()
const maxAnswers = 100000
_, clock, g := start()
now := clock.Now()
// lookups.json lists the answers by client. Here each was last used a
// second before the one listed before it, so the last listed is the
// one used longest ago, and the one dropped.
answers := make([]lookup.Answer, maxAnswers+1)
addr := netip.MustParseAddr("10.0.0.0")
for i := range answers {
answers[i] = lookup.Answer{
Client: netip.PrefixFrom(addr, addr.BitLen()),
Country: germany,
Answered: now,
Used: now.Add(-time.Duration(i) * time.Second),
}
addr = addr.Next()
}
g.Load(answers)
got := g.Snapshot()
if len(got) != maxAnswers || got[0] != answers[0] ||
got[maxAnswers-1] != answers[maxAnswers-1] {
t.Errorf("%d answers kept, from %s to %s; want %d, from %s to %s",
len(got), got[0].Client, got[len(got)-1].Client, maxAnswers,
answers[0].Client, answers[maxAnswers-1].Client)
}
}
+116
View File
@@ -0,0 +1,116 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label under which the countries outside the busiest are
// counted.
const other = "other"
// countries are the metrics by the client's country, for requests whose
// client's country is known. The topN busiest countries, by their requests
// since the start, have series of their own, and the others are counted
// under other, so that there are never more than topN + 1 series. A
// country that drops out of the busiest loses its series, and its next
// requests are counted under other; one that becomes one of them gets a
// series that counts from then on. Each series therefore only ever goes
// up.
type countries struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are the requests the country lists refused.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's requests since the start, by which the
// countries are ranked. GeoJS gives two-letter codes, so it holds at
// most a few hundred.
seen map[string]int64
// top are the countries with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *countries {
byCountry := []string{"country"}
return &countries{
topN: topN,
requests: counterVec("smallwebwaf_country_requests_total",
"Requests, by the client's country.", byCountry),
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
"Request body bytes, by the client's country.", byCountry),
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
"Response body bytes, by the client's country.", byCountry),
refused: counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
byCountry),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// add counts a request from its log line, whose country is known.
func (c *countries) add(line *requestlog.Line) {
c.mu.Lock()
defer c.mu.Unlock()
c.seen[line.Country]++
label := c.label(line.Country)
c.requests.WithLabelValues(label).Inc()
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if line.Action == requestlog.ActionCountryDenied {
c.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from country is counted under: the
// country while it is one of the busiest, other while it is not. A
// country busier than the least busy of them takes its place, and that
// country's series are dropped.
func (c *countries) label(country string) string {
if c.top[country] {
return country
}
if len(c.top) < c.topN {
c.top[country] = true
return country
}
least := ""
for top := range c.top {
if least == "" || c.seen[top] < c.seen[least] {
least = top
}
}
if c.seen[country] <= c.seen[least] {
return other
}
delete(c.top, least)
for _, vec := range []*prometheus.CounterVec{
c.requests, c.requestBytes, c.responseBytes, c.refused,
} {
vec.DeleteLabelValues(least)
}
c.top[country] = true
return country
}
+380
View File
@@ -0,0 +1,380 @@
// Package metrics keeps smallwebwaf's Prometheus metrics, as the "Metrics
// endpoint" section of SPEC.md lists them, and serves them in the
// Prometheus text format. No metric carries a client's address.
package metrics
import (
"net/http"
"strconv"
"time"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/collectors"
"github.com/prometheus/client_golang/prometheus/promhttp"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
type Metrics struct {
registry *prometheus.Registry
handler http.Handler
inFlight prometheus.Gauge
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
requestDuration prometheus.Histogram
upstreamDuration prometheus.Histogram
rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec
// ruleMatches are made by AddRules.
ruleMatches *prometheus.CounterVec
countries *countries
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
// that failed. GeoJSUnanswered are the requests whose client counted
// as coming from an unknown country because GeoJS had not answered
// about it in time.
GeoJSRequests prometheus.Counter
GeoJSFailures prometheus.Counter
GeoJSUnanswered prometheus.Counter
stateFileWrites *prometheus.CounterVec
stateFileWriteFailures *prometheus.CounterVec
stateFileLastWrite *prometheus.GaugeVec
stateFileSize *prometheus.GaugeVec
stateFileEditsTakenIn *prometheus.CounterVec
stateFileEditsSetAside *prometheus.CounterVec
}
// New returns the metrics, with the Go runtime's and the process's own.
// topN is how many countries get series of their own
// (SWWAF_METRICS_TOP_N).
func New(topN int) *Metrics {
byStatus := []string{"status_class", "action"}
byFile := []string{"file"}
m := &Metrics{
registry: prometheus.NewRegistry(),
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
Name: "smallwebwaf_requests_in_flight",
Help: "Requests under way.",
}),
requests: counterVec("smallwebwaf_requests_total",
"Requests, by the class of their status and their action.", byStatus),
requestBytes: counterVec("smallwebwaf_request_bytes_total",
"Request body bytes, by the class of the status and the action.",
byStatus),
responseBytes: counterVec("smallwebwaf_response_bytes_total",
"Response body bytes, by the class of the status and the action.",
byStatus),
requestDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
Name: "smallwebwaf_request_duration_seconds",
Help: "How long requests took, from their arrival to their end.",
}),
upstreamDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
Name: "smallwebwaf_upstream_duration_seconds",
Help: "How long requests passed to the app took, from then to their end.",
}),
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
"Requests that broke a rate limit, by its window.",
[]string{"window"}),
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
"Requests that passed a size or time limit, by its setting.",
[]string{"limit"}),
offences: counterVec("smallwebwaf_offences_total",
"Offences, by kind.", []string{"kind"}),
countries: newCountries(topN),
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_requests_total",
Help: "Requests to GeoJS.",
}),
GeoJSFailures: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_failures_total",
Help: "Requests to GeoJS that failed.",
}),
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_unanswered_total",
Help: "Requests whose client counted as coming from an unknown " +
"country because GeoJS had not answered about it in time.",
}),
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
"Writes of each state file.", byFile),
stateFileWriteFailures: counterVec("smallwebwaf_state_file_write_failures_total",
"Writes of each state file that failed.", byFile),
stateFileLastWrite: gaugeVec("smallwebwaf_state_file_last_write_timestamp_seconds",
"When each state file was last written, in seconds since 1970.", byFile),
stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes",
"The size of each state file, as it was last written.", byFile),
stateFileEditsTakenIn: counterVec("smallwebwaf_state_file_edits_taken_in_total",
"Edits of each state file taken in while running.", byFile),
stateFileEditsSetAside: counterVec("smallwebwaf_state_file_edits_set_aside_total",
"Edits of each state file renamed to <name>.bad because they did not parse.",
byFile),
}
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
m.registry.MustRegister(
collectors.NewGoCollector(),
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
m.inFlight, m.requests, m.requestBytes, m.responseBytes,
m.requestDuration, m.upstreamDuration,
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
m.countries.refused,
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize,
m.stateFileEditsTakenIn, m.stateFileEditsSetAside,
)
return m
}
// AddBansAndClients adds the metrics read from the ledger and the table
// of clients as the metrics are asked for: the bans made since the start,
// by cause, the bans active and permanent at now, and the clients in the
// table.
func (m *Metrics) AddBansAndClients(
ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time,
) {
for _, cause := range []string{bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin} {
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_bans_made_total",
Help: "Bans made, by cause.",
ConstLabels: prometheus.Labels{"cause": cause},
}, func() float64 {
return float64(ledger.Made(cause))
}))
}
m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_active_bans",
Help: "Bans active now, the permanent ones included.",
}, func() float64 {
active, _ := ledger.Count(now())
return float64(active)
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_permanent_bans",
Help: "Permanent bans not lifted.",
}, func() float64 {
_, permanent := ledger.Count(now())
return float64(permanent)
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_tracked_clients",
Help: "Clients in the table of clients.",
}, func() float64 {
return float64(limiter.Len())
}),
)
}
// AddRules adds the metrics of the rule files: the requests that matched
// each rule, which RuleMatched counts, and the rules loaded from
// ruleFiles, read as the metrics are asked for. It is called once, before
// RuleMatched.
func (m *Metrics) AddRules(ruleFiles *rules.Files) {
m.ruleMatches = counterVec("smallwebwaf_rule_matches_total",
"Requests that matched a rule of the rule files, by its id and action.",
[]string{"rule_id", "action"})
m.registry.MustRegister(m.ruleMatches,
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_rules_loaded",
Help: "Rules loaded from the rule files.",
}, func() float64 {
return float64(ruleFiles.Len())
}))
}
// AddRemoteLog adds the metrics of sending the log lines to
// SWWAF_LOG_REMOTE_URL, read from remote as the metrics are asked for: the
// lines sent, those dropped, and those waiting in the buffer.
func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_remote_log_lines_sent_total",
Help: "Log lines sent to SWWAF_LOG_REMOTE_URL.",
}, func() float64 {
return float64(remote.Sent())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_remote_log_lines_dropped_total",
Help: "Log lines dropped: the oldest in a full buffer, and those " +
"whose sending failed.",
}, func() float64 {
return float64(remote.Dropped())
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_remote_log_buffer_depth",
Help: "Log lines in the buffer, waiting to be sent.",
}, func() float64 {
return float64(remote.Depth())
}),
)
}
// AddAlerts adds the metrics of the alerts sent to each destination set,
// read from queue as the metrics are asked for, by destination: the
// alerts sent, the requests to the destination that failed, the alerts
// held back, which are the same for every destination, and those
// dropped. With no destination set, it adds none.
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
for _, name := range queue.DestinationsSet() {
destination := prometheus.Labels{"destination": name}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_sent_total",
Help: "Alerts the destination took.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Sent)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_failed_total",
Help: "Requests to the destination that failed.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Failed)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_suppressed_total",
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Suppressed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_dropped_total",
Help: "Alerts dropped, the oldest first, from a full queue, and " +
"alerts given up as the destination refused them.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Dropped)
}),
)
}
}
// ServeHTTP answers with the metrics in the Prometheus text format.
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
m.handler.ServeHTTP(w, r)
}
// RequestStarted counts a request as under way.
func (m *Metrics) RequestStarted() {
m.inFlight.Inc()
}
// RequestEnded counts a request that has ended, from its log line. limit
// is the setting whose size or time limit the request passed, "" if none.
// duration is how long the request took, and upstreamDuration how long it
// took from when it was passed to the app, zero if it was not.
func (m *Metrics) RequestEnded(
line *requestlog.Line, limit string, duration, upstreamDuration time.Duration,
) {
m.inFlight.Dec()
class := statusClass(line.Status)
m.requests.WithLabelValues(class, line.Action).Inc()
m.requestBytes.WithLabelValues(class, line.Action).Add(float64(line.RequestBytes))
m.responseBytes.WithLabelValues(class, line.Action).Add(float64(line.ResponseBytes))
m.requestDuration.Observe(duration.Seconds())
if upstreamDuration > 0 {
m.upstreamDuration.Observe(upstreamDuration.Seconds())
}
if line.LimitHit != "" {
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
}
if limit != "" {
m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc()
}
if line.Offence != "" {
m.offences.WithLabelValues(line.Offence).Inc()
}
if line.Country != "" {
m.countries.add(line)
}
}
// RuleMatched counts a request that matched the rule id, whose action is
// action.
func (m *Metrics) RuleMatched(id, action string) {
m.ruleMatches.WithLabelValues(id, action).Inc()
}
// StateFileWritten counts a write of the state file name, of size bytes,
// that ended with err.
func (m *Metrics) StateFileWritten(name string, size int, err error) {
m.stateFileWrites.WithLabelValues(name).Inc()
// The series of failures is there from the first write, at zero until
// one fails.
failures := m.stateFileWriteFailures.WithLabelValues(name)
if err != nil {
failures.Inc()
return
}
m.stateFileLastWrite.WithLabelValues(name).SetToCurrentTime()
m.stateFileSize.WithLabelValues(name).Set(float64(size))
}
// StateFileEditTakenIn counts an admin's edit of the state file name
// taken in while smallwebwaf runs.
func (m *Metrics) StateFileEditTakenIn(name string) {
m.stateFileEditsTakenIn.WithLabelValues(name).Inc()
}
// StateFileEditSetAside counts an admin's edit of the state file name
// renamed to name.bad because it did not parse.
func (m *Metrics) StateFileEditSetAside(name string) {
m.stateFileEditsSetAside.WithLabelValues(name).Inc()
}
// statusClass returns the class of status, such as 2xx, or none when no
// status was sent.
func statusClass(status int) string {
if status == 0 {
return "none"
}
// A status's class is its hundreds: 404 is in 4xx.
const hundred = 100
return strconv.Itoa(status/hundred) + "xx"
}
// counterVec returns a counter named name, described by help, with a
// series for each set of values of labels.
func counterVec(name, help string, labels []string) *prometheus.CounterVec {
return prometheus.NewCounterVec(prometheus.CounterOpts{Name: name, Help: help},
labels)
}
// gaugeVec returns a gauge named name, described by help, with a series
// for each set of values of labels.
func gaugeVec(name, help string, labels []string) *prometheus.GaugeVec {
return prometheus.NewGaugeVec(prometheus.GaugeOpts{Name: name, Help: help}, labels)
}
+325
View File
@@ -0,0 +1,325 @@
package proxy
import (
"bytes"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/netip"
"os"
"strings"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
// banBodyMaxBytes is the most of the body of a request to add a ban that
// is read; its three fields need far less.
const banBodyMaxBytes = 4 << 10
// permanent is how the log line and the ban endpoint name a ban that
// never ends.
const permanent = "permanent"
var (
errNotBanToAdd = errors.New(
"the body is not a JSON object of netblock, duration and reason")
errNotNetblock = errors.New(
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
errMappedNetblock = errors.New(
"is IPv4-mapped: give the IPv4 netblock, such as 203.0.113.0/24")
errZone = errors.New("has a zone, which a netblock cannot have")
errNotDuration = errors.New(
"is not a duration above zero, such as 1h or 7d, or permanent")
errNotAddress = errors.New("is not an address, such as 203.0.113.9")
)
// answerAdmin answers a request for smallwebwaf itself, under
// /_smallwebwaf/, once it has passed the checks. Each endpoint needs a
// token, sent as Authorization: Bearer <token>: the metrics
// SWWAF_METRICS_TOKEN, the others SWWAF_ADMIN_TOKEN. A request without
// it is refused with 401. An endpoint whose token is unset answers 404,
// as any other request under /_smallwebwaf/ does.
func (rq *request) answerAdmin() {
rq.line.Action = requestlog.ActionAdmin
rq.startClientResponseTimeout()
token, answer := rq.endpoint()
switch {
case token == "":
http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound)
case !hasToken(rq.in, token):
rq.out.Header().Set("WWW-Authenticate", "Bearer")
rq.answer(refusal{
status: http.StatusUnauthorized,
action: requestlog.ActionAdmin,
})
default:
answer()
}
}
// endpoint returns the token the request's endpoint needs, and what
// answers the request there; "" when there is no such endpoint.
func (rq *request) endpoint() (string, func()) {
cfg := rq.h.config
method, path := rq.in.Method, rq.in.URL.Path
switch {
case method == http.MethodGet && path == MetricsPath:
return cfg.MetricsToken, func() { rq.h.metrics.ServeHTTP(rq.out, rq.in) }
case method == http.MethodGet && path == BansPath:
return cfg.AdminToken, rq.listBans
case method == http.MethodPost && path == BansPath:
return cfg.AdminToken, rq.addBan
case method == http.MethodDelete && strings.HasPrefix(path, BansPath+"/"):
return cfg.AdminToken, rq.liftBans
case method == http.MethodGet && strings.HasPrefix(path, ClientsPath):
return cfg.AdminToken, rq.showClient
default:
return "", nil
}
}
// hasToken reports whether r carries token, as Authorization: Bearer
// <token>.
func hasToken(r *http.Request, token string) bool {
scheme, sent, _ := strings.Cut(r.Header.Get("Authorization"), " ")
return strings.EqualFold(scheme, "Bearer") &&
subtle.ConstantTimeCompare([]byte(sent), []byte(token)) == 1
}
// listBans answers GET BansPath with every ban held.
func (rq *request) listBans() {
rq.answerBans(rq.h.ledger.Snapshot())
}
// banToAdd is the body of POST BansPath.
type banToAdd struct {
// Netblock is a netblock, or a client's address, which stands for the
// netblock a ban on that client covers.
Netblock string `json:"netblock"`
// Duration is how long the ban lasts, as a setting gives a duration,
// or permanent.
Duration string `json:"duration"`
Reason string `json:"reason"`
}
// addBan answers POST BansPath: it bans the netblock the body names, as
// an admin, from now for the duration the body gives, with its reason,
// and answers with that ban.
func (rq *request) addBan() {
// The body must arrive within SWWAF_CLIENT_REQUEST_TIMEOUT, as any
// other request's must.
rq.stopReadingBody(rq.clientRequestDeadline())
toAdd, err := rq.readBanToAdd()
if refused := rq.refused.Load(); refused != nil {
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
return
}
if errors.Is(err, os.ErrDeadlineExceeded) {
rq.answer(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
})
return
}
var (
netblock netip.Prefix
expires time.Time
now = rq.h.now()
)
if err == nil {
netblock, err = rq.h.banNetblock(toAdd.Netblock)
}
if err == nil {
expires, err = expiry(toAdd.Duration, now)
}
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
ban := rq.h.ledger.BanForAdmin(netblock, now, expires, toAdd.Reason)
rq.answerBans([]bans.Ban{ban})
}
// readBanToAdd reads the body of POST BansPath: a JSON object with
// nothing but whitespace after it, in at most banBodyMaxBytes.
func (rq *request) readBanToAdd() (banToAdd, error) {
var body io.ReadCloser = http.NoBody
if rq.body != nil {
body = rq.body
}
data, err := io.ReadAll(http.MaxBytesReader(nil, body, banBodyMaxBytes))
if err != nil {
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
}
var toAdd banToAdd
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
err = decoder.Decode(&toAdd)
if err != nil {
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
}
// Token returns io.EOF only when nothing but whitespace is left.
_, err = decoder.Token()
if !errors.Is(err, io.EOF) {
return banToAdd{}, fmt.Errorf("%w: more follows the object", errNotBanToAdd)
}
return toAdd, nil
}
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a
// client's address, which stands for the netblock a ban on that client
// covers. An IPv4-mapped netblock, such as ::ffff:203.0.113.0/120, is
// refused, since a client's address is looked up as IPv4 and a ban on it
// would refuse nothing, and so is a value with a zone.
func (h *handler) banNetblock(value string) (netip.Prefix, error) {
netblock, err := netip.ParsePrefix(value)
if err == nil {
if netblock.Addr().Is4In6() {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errMappedNetblock)
}
return netblock, nil
}
// ParsePrefix refuses a zone, but ParseAddr reads the /48 of
// 2001:db8::1%x/48 as part of the zone.
addr, err := netip.ParseAddr(value)
if err != nil {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
}
if addr.Zone() != "" {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errZone)
}
return h.netblock(addr), nil
}
// expiry returns when a ban made at now for duration ends: duration
// later, for a duration as a setting gives one, or zero for permanent.
func expiry(duration string, now time.Time) (time.Time, error) {
if duration == permanent {
return time.Time{}, nil
}
length, err := config.ParseDurationNotOff(duration)
if err != nil {
return time.Time{}, fmt.Errorf("duration %q %w", duration, errNotDuration)
}
return now.Add(length), nil
}
// liftBans answers DELETE BansPath/<client>: it lifts every ban active on
// a netblock the client's address is in, and answers with those bans, or
// with 404 when none is active.
func (rq *request) liftBans() {
client, err := pathAddress(rq.in.URL.Path, BansPath+"/")
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
lifted := rq.h.ledger.Lift(client, rq.h.now())
if len(lifted) == 0 {
http.Error(rq.out, "no ban is active on "+client.String(), http.StatusNotFound)
return
}
rq.answerBans(lifted)
}
// clientAnswer is the answer to GET ClientsPath<ip>: the client the
// address is, as clients.json holds it, or null when the table of
// clients does not hold it, and the bans on each netblock the address is
// in, as bans.json lists them.
type clientAnswer struct {
Client *ratelimit.Client `json:"client"`
Bans []state.BanEntry `json:"bans"`
}
// showClient answers GET ClientsPath<ip> with what smallwebwaf knows of
// the client: its counters, its history, which holds its country as last
// looked up and its offences, and its bans with their notes.
func (rq *request) showClient() {
addr, err := pathAddress(rq.in.URL.Path, ClientsPath)
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
client, seen := rq.h.limiter.Client(clientGroup(addr))
if seen {
answer.Client = &client
}
rq.answerJSON(answer)
}
// pathAddress reads the client's address that follows prefix in path.
func pathAddress(path, prefix string) (netip.Addr, error) {
value := strings.TrimPrefix(path, prefix)
addr, err := netip.ParseAddr(value)
if err != nil {
return netip.Addr{}, fmt.Errorf("%q %w", value, errNotAddress)
}
return addr.Unmap(), nil
}
// answerBans answers with held under bans, as bans.json lists them.
func (rq *request) answerBans(held []bans.Ban) {
rq.answerJSON(struct {
Bans []state.BanEntry `json:"bans"`
}{state.BanEntries(held)})
}
// answerJSON answers with value as indented JSON.
func (rq *request) answerJSON(value any) {
body, err := json.MarshalIndent(value, "", " ")
if err != nil {
rq.h.processLog.Error("encoding an answer failed", "error", err.Error())
http.Error(rq.out, http.StatusText(http.StatusInternalServerError),
http.StatusInternalServerError)
return
}
rq.out.Header().Set("Content-Type", "application/json")
_, _ = rq.out.Write(append(body, '\n'))
}
+535
View File
@@ -0,0 +1,535 @@
package proxy_test
import (
"encoding/json"
"net/http"
"net/netip"
"slices"
"strconv"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
const (
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set, and adminBearer
// how a request carries it.
adminSecret = "fedcba9876543210fedcba9876543210"
adminBearer = "Bearer " + adminSecret
// adminClient is the client the tests' admin sends its requests from.
adminClient = "192.0.2.10"
// banOtherClient is the body of a request to ban otherClient for an
// hour.
banOtherClient = `{"netblock": "` + otherClient + `", "duration": "1h", ` +
`"reason": "probes for logins"}`
)
func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
t.Parallel()
// The metrics token is set, and opens none of them.
s, clk, server := startWithClock(t, "", map[string]string{metricsToken: token})
server.Ledger.BanForLimit(netip.MustParsePrefix(otherClient+"/32"), clk.Now(),
bans.Notes{})
before := server.Ledger.Snapshot()
// An empty token does not match the unset one either.
for _, authorization := range []string{adminBearer, bearer, "Bearer ", ""} {
for _, e := range adminEndpoints() {
s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
http.StatusNotFound, requestlog.ActionAdmin)
}
}
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before)
}
}
func TestAdminEndpointsNeedTheAdminToken(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
metricsToken: token,
})
// Listing the bans, banning otherClient, lifting that ban, and asking
// about otherClient, in that order. Without the admin token, with the
// metrics token, or with one that differs, each is refused, and
// changes nothing; with the admin token, it is answered.
for _, e := range adminEndpoints() {
before := server.Ledger.Snapshot()
for _, authorization := range []string{
"", bearer, "Bearer " + strings.ToUpper(adminSecret), "Basic " + adminSecret,
} {
got := s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
http.StatusUnauthorized, requestlog.ActionAdmin)
if got.header.Get("WWW-Authenticate") != "Bearer" {
t.Errorf("%s %s with %q was answered without WWW-Authenticate: Bearer",
e.method, e.path, authorization)
}
}
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
e.method, e.path, after, before)
}
got := s.admin(e.method, e.path, e.body, http.StatusOK)
if got.header.Get("Content-Type") != "application/json" {
t.Errorf("%s %s answered %q", e.method, e.path, got.header.Get("Content-Type"))
}
}
// Any other request under /_smallwebwaf/ is not found.
for _, e := range []adminEndpoint{
{http.MethodPut, proxy.BansPath, banOtherClient},
{http.MethodDelete, proxy.BansPath, ""},
{http.MethodGet, proxy.BansPath + "/" + otherClient, ""},
{http.MethodPost, proxy.ClientsPath + otherClient, ""},
{http.MethodGet, strings.TrimSuffix(proxy.ClientsPath, "/"), ""},
} {
s.admin(e.method, e.path, e.body, http.StatusNotFound)
}
}
func TestBanAddedListedAndLiftedThroughTheEndpoints(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
banScopeV4Prefix: "24",
})
// A ban on otherClient bans the /24 a ban on that client covers, so it
// refuses client too, for an hour.
start := clk.Now()
expires := start.Add(time.Hour)
want := state.BanEntry{
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Start: start,
Expires: &expires,
Cause: bans.CauseAdmin,
Reason: "probes for logins",
}
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK),
want)
line := s.get(client, http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != requestlog.FormatTime(expires) {
t.Errorf("the ban ends at %s, want %s", line.BanExpires, expires)
}
// Its notes count the request it refused.
want.Notes.Requests, want.Notes.Refused = 1, 1
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
// Ten minutes on, lifting the bans on client lifts that one, which is
// kept, marked lifted.
clk.advance(10 * time.Minute)
lifted := clk.Now()
want.Lifted = &lifted
wantBans(t, s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "",
http.StatusOK), want)
s.get(client, http.StatusOK, requestlog.ActionForward)
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
// No ban on it is active any more.
s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "", http.StatusNotFound)
}
func TestBanToAddGivesItsNetblockAndDuration(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
banScopeV4Prefix: "24",
})
start := clk.Now()
for _, tc := range []struct {
netblock, duration string
want string
length time.Duration // 0 for a permanent ban
}{
// An address stands for the netblock a ban on that client covers.
{client, "7d", "203.0.113.0/24", 7 * 24 * time.Hour},
{"::ffff:198.51.100.7", "90m", "198.51.100.0/24", 90 * time.Minute},
{"2001:db8:5::1", "permanent", "2001:db8:5::/64", 0},
// A netblock stands for itself, its bits past its length cleared.
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
} {
// Whitespace may follow the object.
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` +
"\r\n"
want := state.BanEntry{
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
}
if tc.length != 0 {
expires := start.Add(tc.length)
want.Expires = &expires
}
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, body, http.StatusOK), want)
}
}
func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{adminToken: adminSecret})
for _, tc := range []struct{ body, want string }{
{"", "the body is not a JSON object of netblock, duration and reason: EOF"},
{"netblock=203.0.113.9", "the body is not a JSON object"},
{
`{"netblock": "203.0.113.9", "duration": "1h", "until": "2027"}`,
`unknown field "until"`,
},
{
`{"netblock": "203.0.113", "duration": "1h"}`,
`netblock "203.0.113" is not an address or a netblock`,
},
// A client's address is looked up as IPv4, so a ban on an
// IPv4-mapped netblock would refuse nothing.
{
`{"netblock": "::ffff:203.0.113.0/120", "duration": "1h"}`,
`netblock "::ffff:203.0.113.0/120" is IPv4-mapped`,
},
// Read as an address, its zone would be "x/48", and its ban on the
// /64 around it.
{
`{"netblock": "2001:db8::1%x/48", "duration": "1h"}`,
`netblock "2001:db8::1%x/48" has a zone`,
},
{
`{"netblock": "fe80::1%eth0", "duration": "1h"}`,
`netblock "fe80::1%eth0" has a zone`,
},
// Anything but whitespace after the object.
{
`{"netblock": "203.0.113.9", "duration": "1h"}` +
`{"netblock": "198.51.100.0/24", "duration": "1h"}`,
"more follows the object",
},
{`{"netblock": "203.0.113.9", "duration": "1h"} x`, "more follows the object"},
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
{
`{"netblock": "203.0.113.9", "duration": "off"}`,
`duration "off" is not a duration above zero`,
},
{
`{"netblock": "203.0.113.9", "duration": "0s"}`,
`duration "0s" is not a duration above zero`,
},
{
`{"netblock": "203.0.113.9", "duration": "forever"}`,
`duration "forever" is not a duration above zero, such as 1h or 7d, ` +
`or permanent`,
},
// Over the 4 KiB read of a body, even when the object comes first.
{
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
strings.Repeat("x", 4<<10) + `"}`,
"request body too large",
},
{
`{"netblock": "203.0.113.9", "duration": "1h"}` + strings.Repeat(" ", 4<<10),
"request body too large",
},
} {
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
if !strings.Contains(string(got.body), tc.want) {
t.Errorf("%.80s was answered %q, want it to say %q", tc.body, got.body, tc.want)
}
}
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestBanToAddOverTheRequestSizeLimitIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
requestMaxBytes: "16",
})
// Sent in a chunk, its length is not announced, so that it is found
// over SWWAF_REQUEST_MAX_BYTES only as it is read.
chunk := `{"netblock": "203.0.113.9", "duration": "1h"}`
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
http.MethodPost, proxy.BansPath,
strconv.FormatInt(int64(len(chunk)), 16)+"\r\n"+chunk+"\r\n0\r\n\r\n",
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestBanToAddSlowerThanTheClientRequestTimeoutIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
metricsToken: token,
clientRequestTimeout: shortTimeoutSetting,
})
// The chunk announces 256 bytes and the rest of it never comes, so only
// the timeout ends the wait. A hold-up of the test process can only
// make the answer later, so the time is checked only for not being
// shorter than the timeout.
start := time.Now()
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
http.MethodPost, proxy.BansPath, "100\r\n"+`{"netblock": "203.0.113.9", `,
http.StatusRequestTimeout, requestlog.ActionTimedOut)
if took := time.Since(start); took < shortTimeout {
t.Errorf("answered after %s, before the timeout of %s ran out", took, shortTimeout)
}
wantLimitHits(t, s.addr, clientRequestTimeout, 1)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
rateLimitPerMinute: "2",
rateLimitExemptNets: adminClient,
})
start := clk.Now()
// Two of otherClient's requests are let through; the third breaks the
// limit of two a minute, and bans it.
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited)
// Asked about by its address in IPv6 form too.
for _, addr := range []string{otherClient, "::ffff:" + otherClient} {
var got struct {
Client *ratelimit.Client `json:"client"`
Bans []state.BanEntry `json:"bans"`
}
decode(t, s.admin(http.MethodGet, proxy.ClientsPath+addr, "", http.StatusOK), &got)
if got.Client == nil {
t.Fatalf("%s: no client", addr)
}
history := got.Client.History
if got.Client.Client != netip.MustParsePrefix(otherClient+"/32") ||
history.Requests != 3 || history.Forwarded != 2 || history.Refused != 1 ||
history.Offences.Limit != 1 || !history.FirstSeen.Equal(start) {
t.Errorf("%s: client %+v", addr, got.Client)
}
if len(got.Bans) != 1 || got.Bans[0].Cause != bans.CauseLimit ||
got.Bans[0].Reason != "requests per minute over the limit of 2" ||
got.Bans[0].Notes.Count != 3 {
t.Errorf("%s: bans %+v, want the one for the broken limit", addr, got.Bans)
}
}
// Of an address no request came from and no ban covers, nothing is
// known.
got := s.admin(http.MethodGet, proxy.ClientsPath+"198.51.100.99", "", http.StatusOK)
if string(got.body) != "{\n \"client\": null,\n \"bans\": []\n}\n" {
t.Errorf("an unknown client is answered\n%s", got.body)
}
s.admin(http.MethodGet, proxy.ClientsPath+"203.0.113", "", http.StatusBadRequest)
s.admin(http.MethodDelete, proxy.BansPath+"/203.0.113.0/24", "",
http.StatusBadRequest)
}
func TestBannedClientIsRefusedAtTheEndpointsEvenWithTheToken(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{adminToken: adminSecret})
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
// otherClient cannot lift its own ban either.
for _, e := range adminEndpoints() {
s.adminRequest(otherClient, adminBearer, e.method, e.path, e.body,
http.StatusForbidden, requestlog.ActionBanned)
}
}
func TestAdminRequestsCountTowardTheLimits(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
rateLimitPerMinute: "2",
})
// A request refused for a missing token and one answered count toward
// the limit of two a minute, so the next breaks it.
s.adminRequest(client, "", http.MethodGet, proxy.BansPath, "",
http.StatusUnauthorized, requestlog.ActionAdmin)
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusOK, requestlog.ActionAdmin)
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusForbidden, requestlog.ActionRateLimited)
}
func TestClientInAllowNetsSkipsTheChecksButNeedsTheToken(t *testing.T) {
t.Parallel()
const allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
s, clk, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
allowNets: allowed,
rateLimitPerMinute: "1",
})
// A ban on it refuses nothing, and its requests are not counted.
server.Ledger.BanForAdmin(netip.MustParsePrefix(allowed+"/32"), clk.Now(),
time.Time{}, "")
for range 2 {
s.adminRequest(allowed, "", http.MethodGet, proxy.BansPath, "",
http.StatusUnauthorized, requestlog.ActionAdmin)
s.adminRequest(allowed, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusOK, requestlog.ActionAdmin)
}
}
func TestAdminEndpointsNeedTheTokenInObserveMode(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
mode: observe,
})
for _, e := range adminEndpoints() {
s.adminRequest(adminClient, "", e.method, e.path, e.body,
http.StatusUnauthorized, requestlog.ActionAdmin)
}
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
// adminEndpoint is a request to an endpoint SWWAF_ADMIN_TOKEN opens.
type adminEndpoint struct {
method, path, body string
}
// adminEndpoints returns a request to each endpoint SWWAF_ADMIN_TOKEN
// opens: listing the bans, banning otherClient for an hour, lifting the
// bans on otherClient, and asking about otherClient.
func adminEndpoints() []adminEndpoint {
return []adminEndpoint{
{http.MethodGet, proxy.BansPath, ""},
{http.MethodPost, proxy.BansPath, banOtherClient},
{http.MethodDelete, proxy.BansPath + "/" + otherClient, ""},
{http.MethodGet, proxy.ClientsPath + otherClient, ""},
}
}
// admin sends a request with method for path, with body, from
// adminClient, with the admin token, and checks that it is answered with
// status, its log line's action admin. It returns the answer.
func (s *sender) admin(method, path, body string, status int) answer {
s.t.Helper()
return s.adminRequest(adminClient, adminBearer, method, path, body, status,
requestlog.ActionAdmin)
}
// adminRequest sends a request with method for path, with body, from the
// client at from, with authorization as its Authorization header unless
// it is "", and checks its answer's status and its log line's action, as
// request does. authorization may end in more header lines. A body that
// is not "" has its length announced, unless authorization names
// Transfer-Encoding. It returns the answer.
func (s *sender) adminRequest(
from, authorization, method, path, body string, status int, action string,
) answer {
s.t.Helper()
var header []string
if authorization != "" {
header = append(header, "Authorization: "+authorization)
}
if body != "" && !strings.Contains(authorization, "Transfer-Encoding") {
header = append(header, "Content-Length: "+strconv.Itoa(len(body)))
}
_, got := s.requestWithBody(method, from, path, strings.Join(header, "\r\n"),
body, status, action)
return got
}
// wantBans checks that a ban endpoint answered with want, and no other
// ban.
func wantBans(t *testing.T, got answer, want ...state.BanEntry) {
t.Helper()
var decoded struct {
Bans []state.BanEntry `json:"bans"`
}
decode(t, got, &decoded)
gotJSON, err := json.Marshal(decoded.Bans)
if err != nil {
t.Fatalf("encode %+v: %v", decoded.Bans, err)
}
wantJSON, err := json.Marshal(want)
if err != nil {
t.Fatalf("encode %+v: %v", want, err)
}
if string(gotJSON) != string(wantJSON) {
t.Errorf("bans\n%s\nwant\n%s", gotJSON, wantJSON)
}
}
// decode reads the JSON answer of an endpoint into value.
func decode(t *testing.T, got answer, value any) {
t.Helper()
err := json.Unmarshal(got.body, value)
if err != nil {
t.Fatalf("decode %s: %v", got.body, err)
}
}
+264
View File
@@ -0,0 +1,264 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"reflect"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
// alertInstance is the instance every alert of these tests gives.
alertInstance = "fsn1app1/gitea"
)
func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rateLimitPerMinute: "1",
banScopeV4Prefix: "24",
})
start := clk.Now()
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
netblock := netip.MustParsePrefix("203.0.113.0/24")
ban := server.Ledger.Bans(netblock)[0]
// A request refused under the ban raises no other alert.
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netblock, Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1", Notes: ban.Notes,
}, requestlog.FormatTime(start.Add(time.Hour))))
if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" {
t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes)
}
}
func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
netblock := netip.MustParsePrefix(client + "/32")
other := netip.MustParsePrefix(otherClient + "/32")
// The probe bans the client for seven days, and its next request makes
// the ban permanent. The request after that changes nothing.
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
attackBan := server.Ledger.Bans(netblock)[0]
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
permanentBan := server.Ledger.Bans(netblock)[0]
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
// Another client's probe after its first ban has run out without a
// request makes a permanent ban at once.
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
clk.advance(7 * 24 * time.Hour)
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
otherBans := server.Ledger.Bans(other)
wantAlerts(t, queue,
attackAlert(alerts.EventBan, start, client, attackBan,
requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute), client,
permanentBan, "permanent"),
attackAlert(alerts.EventBan, start.Add(time.Minute), otherClient, otherBans[0],
requestlog.FormatTime(start.Add(time.Minute+7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute+7*24*time.Hour),
otherClient, otherBans[1], "permanent"),
)
}
func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
// A ban for a clear sign of attack, which a request under it would make
// permanent.
group := netip.MustParsePrefix(ipv6Group)
attackBan, _ := server.Ledger.BanForAttack(group, start, bans.Notes{RuleID: "probe"})
// The third request breaks the limit, and so does the fourth, within the
// cooldown, which raises nothing. The probe is a clear sign of attack.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
// No ban is made, and none made permanent.
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan ||
line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+
"for the attack alone, as it was", held, line.BanExpires)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 3 || queue.Suppressed() != 0 {
t.Fatalf("%d alerts wait and %d are held back, want 3 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
limitNotes, _ := waiting[0].Detail["notes"].(bans.Notes)
attackNotes, _ := waiting[1].Detail["notes"].(bans.Notes)
if limitNotes.Limit != 2 || limitNotes.Request.Path != "/" ||
attackNotes.Request.Path != "/.env" {
t.Errorf("the notes are %+v and %+v, want those of the broken limit and "+
"of the probe", limitNotes, attackNotes)
}
// Each alert is the one enforce mode would have raised, with mode
// observe in its detail.
want := []alerts.Alert{
banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 2", Notes: limitNotes,
}, requestlog.FormatTime(start.Add(time.Hour))),
attackAlert(alerts.EventBan, start, otherClient, bans.Ban{
Netblock: netip.MustParsePrefix(otherClient + "/32"), Notes: attackNotes,
}, requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start, ipv6Client, attackBan, permanent),
}
for _, alert := range want {
alert.Detail["mode"] = observe
}
wantAlerts(t, queue, want...)
}
func TestObserveModeWorksOutABanOnlyWhenItsAlertWouldBeSent(t *testing.T) {
t.Parallel()
s, _, _, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
alertMaxPerHour: "2",
})
// The client's third request breaks the limit, and raises the first
// alert of the hour. Its fourth is within the cooldown.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
// The other client's first probe raises the second. Its second probe is
// within the cooldown.
for range 2 {
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
}
// The IPv6 client's third request breaks the limit past the two alerts
// an hour.
for range 3 {
s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
}
// Had the ban been worked out for any of the requests within the
// cooldown or past the two an hour, its alert would have been raised,
// held back and counted.
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || queue.Suppressed() != 0 {
t.Errorf("%d alerts wait and %d are held back, want 2 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
}
// startWithAlerts is startWithClock with alerts to a webhook, which is
// never sent them, and returns the queue they wait in as well.
func startWithAlerts(
t *testing.T, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{
trustedProxies: trustLocalhost,
alertWebhookURL: "https://alerts.example/smallwebwaf",
instanceName: alertInstance,
}
maps.Copy(settings, env)
addr, out, server, queue := startProxyWithAlerts(t, app.URL, "", clk.Now, settings)
return &sender{t: t, addr: addr, out: out}, clk, server, queue
}
// banAlert returns the alert for event, raised by a request from client at
// the time raised, for ban, with its netblock, cause, reason and notes,
// which ends at expires, as the log line gives it.
func banAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return alerts.Alert{
Instance: alertInstance,
Time: raised,
Event: event,
Client: netip.MustParseAddr(client),
Netblock: ban.Netblock,
Reason: ban.Reason,
Detail: map[string]any{
"cause": ban.Cause, "ban_expires": expires, "notes": ban.Notes,
},
}
}
// attackAlert is banAlert for a ban for the probe rule of testRules, with
// the netblock and the notes of ban.
func attackAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return banAlert(event, raised, client, bans.Ban{
Netblock: ban.Netblock, Cause: bans.CauseAttack, Reason: "matched the rule probe",
Notes: ban.Notes,
}, expires)
}
// wantAlerts checks the alerts waiting in queue, in order.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper()
got := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(got) != len(want) {
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
}
for i := range want {
if !reflect.DeepEqual(got[i], want[i]) {
t.Errorf("alert %d is\n%+v\nwant\n%+v", i, got[i], want[i])
}
}
}
+161 -28
View File
@@ -4,8 +4,10 @@ import (
"net/netip" "net/netip"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/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
@@ -14,58 +16,189 @@ func (rq *request) banResponse(action string) *refusal {
return &refusal{status: rq.h.config.BanResponse, action: action} return &refusal{status: rq.h.config.BanResponse, action: action}
} }
// banned reports whether a ban on the client's netblock refuses the // banned reports whether a ban on a netblock the client is in covers the
// request at now, and notes for the log line when that ban ends. // request at now, and notes for the log line when that ban ends. A
// request that makes the ban permanent, or in observe mode would have,
// raises the alert for it.
func (rq *request) banned(now time.Time) bool { func (rq *request) banned(now time.Time) bool {
ban, banned := rq.h.ledger.Check(rq.netblock(), now) check := rq.h.ledger.Check
if rq.h.config.Observe {
check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is
}
ban, banned, madePermanent := check(rq.client, now)
if banned { if banned {
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
} }
if madePermanent {
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
rq.alertBan(ban)
}
return banned return banned
} }
// limitBroken counts the request for the rate limits at now, and reports // limitBroken counts the request for the rate limits at now, notes the
// whether it takes the client over one. Such a request bans the client's // client's counts for the log line, and reports whether the request takes
// netblock, and sets the client's counters back to zero. // the client over a limit. In enforce mode such a request bans the
// client's netblock, and sets the client's counters back to zero; in
// 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) group := clientGroup(rq.client)
hit, over := rq.h.limiter.Count(group, now) counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts
if !over { if !over {
return false return false
} }
ban := rq.h.ledger.BanForLimit(rq.netblock(), now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
},
})
rq.h.limiter.Reset(group)
rq.line.LimitHit = hit.Window rq.line.LimitHit = hit.Window
rq.line.Offence = requestlog.OffenceLimit rq.line.Offence = requestlog.OffenceLimit
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
return true
}
notes := bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return true
}
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
rq.h.limiter.Reset(group)
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
if made {
rq.alertBan(ban)
}
return true return true
} }
// netblock is the netblock a ban on the client covers: its IPv4 address, // 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,
// and raises the alert for the ban it would have made, if that alert
// would be sent.
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
return
}
notes := bans.Notes{
Country: rq.line.Country,
RuleID: rule.ID,
Target: rule.Target,
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return
}
ban, made := rq.h.ledger.BanForAttack(netblock, now, notes)
rq.line.BanExpires = banExpires(ban)
if made {
rq.alertBan(ban)
}
}
// wouldAlertBan reports whether the alert for a ban on netblock for cause
// made at now would be sent. In observe mode the ban the request would
// have made is worked out only then, at most once per
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
// netblock's requests, which can mean going through every client.
func (rq *request) wouldAlertBan(
netblock netip.Prefix, now time.Time, cause string,
) bool {
event := alerts.EventBan
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
event = alerts.EventPermanentBan
}
return rq.h.alerts.WouldSend(event, netblock)
}
// alertBan raises the alert for ban, which the request made, or made
// permanent: permanent_ban for a permanent ban, ban for another. Its
// detail gives the ban's cause, when it ends, and its notes, and in
// observe mode, where ban is the ban that would have been made, or made
// permanent, mode, observe.
func (rq *request) alertBan(ban bans.Ban) {
event := alerts.EventBan
if ban.Permanent() {
event = alerts.EventPermanentBan
}
detail := map[string]any{
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
}
if rq.h.config.Observe {
detail["mode"] = "observe"
}
rq.h.alerts.Raise(alerts.Alert{
Event: event,
Client: rq.client,
Netblock: ban.Netblock,
Country: ban.Notes.Country,
Reason: ban.Reason,
Detail: detail,
})
}
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, or in
// observe mode as it would have been, as the notes of the ban it makes
// keep it.
func (rq *request) noted(now time.Time) bans.Request {
return bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
}
}
// netblockRequests is how many requests netblock has sent since it was
// first seen, this one included: the histories count it only once it has
// ended.
func (rq *request) netblockRequests(netblock netip.Prefix) int64 {
return rq.h.limiter.Requests(netblock) + 1
}
// netblock is the netblock a ban on client covers: its IPv4 address,
// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup // widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup
// counts it in. // counts it in.
func (rq *request) netblock() netip.Prefix { func (h *handler) netblock(client netip.Addr) netip.Prefix {
addr := rq.client.Unmap() addr := client.Unmap()
if addr.Is4() { if addr.Is4() {
return netip.PrefixFrom(addr, rq.h.config.BanScopeV4Prefix).Masked() return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked()
} }
return clientGroup(addr) return clientGroup(addr)
@@ -75,7 +208,7 @@ func (rq *request) netblock() netip.Prefix {
// permanent. // permanent.
func banExpires(ban bans.Ban) string { func banExpires(ban bans.Ban) string {
if ban.Permanent() { if ban.Permanent() {
return "permanent" return permanent
} }
return requestlog.FormatTime(ban.Expires) return requestlog.FormatTime(ban.Expires)
+51 -14
View File
@@ -133,7 +133,7 @@ func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) {
s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
banned := proxy.LedgerOf(server).Bans(netip.MustParsePrefix(client + "/32")) banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(banned) != 2 || banned[0].Notes.Refused != 3 { if len(banned) != 2 || banned[0].Notes.Refused != 3 {
t.Errorf("bans %+v, want two, the first with 3 requests refused", banned) t.Errorf("bans %+v, want two, the first with 3 requests refused", banned)
} }
@@ -278,6 +278,8 @@ func TestBanNotes(t *testing.T) {
Netblock: netblock, Netblock: netblock,
Start: start, Start: start,
Expires: start.Add(time.Hour), Expires: start.Add(time.Hour),
Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1",
Notes: bans.Notes{ Notes: bans.Notes{
Country: "DE", Country: "DE",
Limit: 1, Limit: 1,
@@ -291,12 +293,15 @@ func TestBanNotes(t *testing.T) {
Status: http.StatusForbidden, Status: http.StatusForbidden,
UserAgent: userAgent, UserAgent: userAgent,
}, },
// The one let through, the one that broke the limit and the two
// refused under the ban.
Requests: 4,
Refused: 2, Refused: 2,
EarlierBans: 0, EarlierBans: bans.EarlierBans{},
}, },
} }
ledger := proxy.LedgerOf(server) ledger := server.Ledger
got := ledger.Bans(netblock) got := ledger.Bans(netblock)
if len(got) != 1 || got[0] != want { if len(got) != 1 || got[0] != want {
@@ -309,8 +314,8 @@ func TestBanNotes(t *testing.T) {
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
got = ledger.Bans(netblock) got = ledger.Bans(netblock)
if len(got) != 2 || got[1].Notes.EarlierBans != 1 { if len(got) != 2 || got[1].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("bans %+v, want two, the second with one earlier ban", got) t.Errorf("bans %+v, want two, the second with one earlier ban for a limit", got)
} }
} }
@@ -361,7 +366,7 @@ func (c *clock) advance(d time.Duration) {
// set to midnight, the start of a bucket in every window. // set to midnight, the start of a bucket in every window.
func startWithClock( func startWithClock(
t *testing.T, geojsURL string, env map[string]string, t *testing.T, geojsURL string, env map[string]string,
) (*sender, *clock, *http.Server) { ) (*sender, *clock, *proxy.Server) {
t.Helper() t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {}) app := startApp(t, func(http.ResponseWriter, *http.Request) {})
@@ -399,36 +404,68 @@ func (s *sender) get(from string, status int, action string) logLine {
func (s *sender) request(from, path string, status int, action string) logLine { func (s *sender) request(from, path string, status int, action string) logLine {
s.t.Helper() s.t.Helper()
line, _ := s.requestWithHeader(from, path, "", status, action)
return line
}
// requestWithHeader is request with header, such as "Authorization:
// Bearer x", added to the request unless it is "". It returns the body of
// the answer too.
func (s *sender) requestWithHeader(
from, path, header string, status int, action string,
) (logLine, string) {
s.t.Helper()
line, got := s.requestWithBody(http.MethodGet, from, path, header, "", status, action)
return line, string(got.body)
}
// requestWithBody is requestWithHeader for a request with method, whose
// body is sent as it is after the headers, header holding its
// Content-Length or Transfer-Encoding. header may hold several lines,
// separated by "\r\n". It returns the whole answer.
func (s *sender) requestWithBody(
method, from, path, header, body string, status int, action string,
) (logLine, answer) {
s.t.Helper()
if header != "" {
header += "\r\n"
}
conn := dial(s.t, s.addr) conn := dial(s.t, s.addr)
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+ send(s.t, conn, method+" "+path+" HTTP/1.1\r\nHost: "+appHost+
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n\r\n") "\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+
header+"\r\n"+body)
err := conn.SetReadDeadline(time.Now().Add(waitLimit)) err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil { if err != nil {
s.t.Fatalf("set read deadline: %v", err) s.t.Fatalf("set read deadline: %v", err)
} }
got := 0 var got answer
res, err := http.ReadResponse(bufio.NewReader(conn), nil) res, err := http.ReadResponse(bufio.NewReader(conn), nil)
switch { switch {
case err == nil: case err == nil:
got = readAnswer(res).status got = readAnswer(res)
case !errors.Is(err, io.ErrUnexpectedEOF): case !errors.Is(err, io.ErrUnexpectedEOF):
s.t.Fatalf("read response: %v", err) s.t.Fatalf("read response: %v", err)
} }
_ = conn.Close() _ = conn.Close()
if got != status { if got.status != status {
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, got, s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from,
status) got.status, status)
} }
line := s.out.requestLines(s.t, s.sent+1)[s.sent] line := s.out.requestLines(s.t, s.sent+1)[s.sent]
s.sent++ s.sent++
wantLine(s.t, line, status, action) wantLine(s.t, line, status, action)
return line return line, got
} }
+2
View File
@@ -46,6 +46,7 @@ func (b *requestBody) Read(p []byte) (int, error) {
b.rq.refuse(refusal{ b.rq.refuse(refusal{
status: http.StatusRequestEntityTooLarge, status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge, action: requestlog.ActionTooLarge,
limit: "SWWAF_REQUEST_MAX_BYTES",
}) })
} }
@@ -81,6 +82,7 @@ func (b *responseBody) Read(p []byte) (int, error) {
b.rq.refuse(refusal{ b.rq.refuse(refusal{
status: http.StatusBadGateway, status: http.StatusBadGateway,
action: requestlog.ActionTooLarge, action: requestlog.ActionTooLarge,
limit: "SWWAF_RESPONSE_MAX_BYTES",
}) })
return n, errResponseTooLarge return n, errResponseTooLarge
+28
View File
@@ -1,6 +1,7 @@
package proxy package proxy
import ( import (
"crypto/rand"
"net/http" "net/http"
"net/netip" "net/netip"
"slices" "slices"
@@ -48,6 +49,33 @@ func clientAddress(
return client return client
} }
// requestIDHeader carries the request's id, from traefik and to the app.
const requestIDHeader = "X-Request-ID"
// requestID is the request's id: the one a trusted proxy sent, or a new
// random one. A peer outside the trusted proxies did not come through
// traefik, so the id it sends is its own claim, and is replaced.
func requestID(r *http.Request, peerTrusted bool) string {
id := r.Header.Get(requestIDHeader)
if !peerTrusted || id == "" {
id = rand.Text()
}
return id
}
// scheme is how the client reached traefik, as a trusted proxy says in
// X-Forwarded-Proto, or otherwise http, the only scheme smallwebwaf
// serves.
func scheme(r *http.Request, peerTrusted bool) string {
proto := r.Header.Get("X-Forwarded-Proto")
if !peerTrusted || proto == "" {
return "http"
}
return proto
}
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client. // ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
const ipv6GroupPrefix = 64 const ipv6GroupPrefix = 64
+17 -13
View File
@@ -14,10 +14,14 @@ const (
appHost = "app.example" appHost = "app.example"
// client is the client's address, as a proxy names it. // client is the client's address, as a proxy names it.
client = "203.0.113.9" client = "203.0.113.9"
// forwardedFor is the header that lists the client and its proxies. // forwardedFor is the header that lists the client and its proxies,
forwardedFor = "X-Forwarded-For" // and forwardedProto the one that gives the scheme the client used.
// secure is the scheme a client reached traefik with. forwardedFor = "X-Forwarded-For"
forwardedProto = "X-Forwarded-Proto"
// secure is the scheme a client reached traefik with, and plain the
// one smallwebwaf serves.
secure = "https" secure = "https"
plain = "http"
) )
// appHeaders is what the app tells about the headers it received. // appHeaders is what the app tells about the headers it received.
@@ -65,13 +69,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
func clientAddressCases() []clientAddressCase { func clientAddressCases() []clientAddressCase {
trusted := map[string]string{trustedProxies: trustLocalhost} trusted := map[string]string{trustedProxies: trustLocalhost}
forged := http.Header{ forged := http.Header{
forwardedFor: {client}, forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"}, "X-Forwarded-Host": {"forged.example"},
"X-Forwarded-Proto": {secure}, forwardedProto: {secure},
"X-Real-Ip": {client}, "X-Real-Ip": {client},
} }
replaced := appHeaders{ replaced := appHeaders{
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http", ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
} }
return []clientAddressCase{{ return []clientAddressCase{{
@@ -87,10 +91,10 @@ func clientAddressCases() []clientAddressCase {
"outside the trusted proxies from the right", "outside the trusted proxies from the right",
env: trusted, env: trusted,
header: http.Header{ header: http.Header{
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"}, forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost}, "X-Forwarded-Host": {appHost},
"X-Forwarded-Proto": {secure}, forwardedProto: {secure},
"X-Real-Ip": {client}, "X-Real-Ip": {client},
}, },
wantClient: client, wantClient: client,
wantApp: appHeaders{ wantApp: appHeaders{
@@ -138,7 +142,7 @@ func requestWithHeaders(
Host: r.Host, Host: r.Host,
ForwardedFor: r.Header.Get(forwardedFor), ForwardedFor: r.Header.Get(forwardedFor),
ForwardedHost: r.Header.Get("X-Forwarded-Host"), ForwardedHost: r.Header.Get("X-Forwarded-Host"),
ForwardedProto: r.Header.Get("X-Forwarded-Proto"), ForwardedProto: r.Header.Get(forwardedProto),
RealIP: r.Header.Get("X-Real-IP"), RealIP: r.Header.Get("X-Real-IP"),
}) })
}) })
-15
View File
@@ -1,15 +0,0 @@
package proxy
import (
"net/http"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
// LedgerOf returns the ban ledger of a server New returned, so that the
// tests can read the bans' notes.
func LedgerOf(server *http.Server) *bans.Ledger {
h, _ := server.Handler.(*handler)
return h.ledger
}
+12 -3
View File
@@ -21,14 +21,18 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
// the last one would have it refused. // the last one would have it refused.
addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"}) addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"})
const healthChecks = 3 const (
healthChecks = 3
contentType = "text/plain; charset=utf-8"
)
for range healthChecks { for range healthChecks {
got := get(t, addr, proxy.HealthPath) got := get(t, addr, proxy.HealthPath)
wantStatus(t, got, http.StatusOK) wantStatus(t, got, http.StatusOK)
if string(got.body) != "ok\n" { if string(got.body) != "ok\n" || got.header.Get("Content-Type") != contentType {
t.Errorf("health endpoint answered %q, want ok", got.body) t.Errorf("health endpoint answered %q with Content-Type %q, want ok "+
"with %q", got.body, got.header.Get("Content-Type"), contentType)
} }
} }
@@ -37,6 +41,11 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
lines := out.requestLines(t, healthChecks+1) lines := out.requestLines(t, healthChecks+1)
for _, line := range lines[:healthChecks] { for _, line := range lines[:healthChecks] {
wantLine(t, line, http.StatusOK, requestlog.ActionAdmin) wantLine(t, line, http.StatusOK, requestlog.ActionAdmin)
if line.ResponseContentType != contentType {
t.Errorf("health check's log line has response_content_type %q, "+
"want %q", line.ResponseContentType, contentType)
}
} }
wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward) wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward)
+124
View File
@@ -0,0 +1,124 @@
package proxy_test
import (
"io"
"net/http"
"net/netip"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "2",
deniedCountries: "kp",
})
start := clk.Now()
// Two let through, one over the limit, which bans the client, and one
// refused under that ban, for which the country is not looked up.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
clk.advance(time.Second)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
clk.advance(time.Second)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
want := ratelimit.History{
FirstSeen: start,
LastSeen: start.Add(2 * time.Second),
Country: "DE",
LookedUp: start.Add(time.Second),
Requests: 4,
Forwarded: 2,
Refused: 2,
// The app answers with no body, smallwebwaf with its status text.
ResponseBytes: 2 * int64(len("Forbidden\n")),
Responses: ratelimit.Responses{Status2xx: 2, Status4xx: 2},
Offences: ratelimit.Offences{Limit: 1},
}
got := historyOf(t, server, fromDE)
if got != want {
t.Errorf("history\n%+v\nwant\n%+v", got, want)
}
}
func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
_, _ = io.WriteString(w, "hello")
})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
wantStatus(t, got, http.StatusOK)
out.requestLine(t)
history := historyOf(t, server, localhost)
if history.RequestBytes != 3 || history.ResponseBytes != 5 {
t.Errorf("history counts %d bytes in and %d out, want 3 and 5",
history.RequestBytes, history.ResponseBytes)
}
}
func TestHealthEndpointIsNotInTheHistory(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
out.requestLine(t)
if clients := server.Limiter.Snapshot(); len(clients) != 0 {
t.Errorf("the table holds %+v, want no client", clients)
}
}
func TestRequestForSmallwebwafIsRefusedOnlyWithoutTheToken(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now,
map[string]string{metricsToken: token})
// The metrics and the 404 are neither forwarded nor refused; the 401
// is refused.
scrape(t, addr)
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
wantStatus(t, get(t, addr, proxy.MetricsPath), http.StatusUnauthorized)
out.requestLines(t, 3)
history := historyOf(t, server, localhost)
if history.Requests != 3 || history.Forwarded != 0 || history.Refused != 1 {
t.Errorf("history counts %d requests, %d forwarded and %d refused, "+
"want 3, 0 and 1", history.Requests, history.Forwarded, history.Refused)
}
}
// historyOf returns the history of the client at addr.
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
t.Helper()
client := netip.MustParsePrefix(addr + "/32")
for _, c := range server.Limiter.Snapshot() {
if c.Client == client {
return c.History
}
}
t.Fatalf("%s is not in the table", client)
return ratelimit.History{}
}
+16
View File
@@ -55,6 +55,7 @@ func TestRequestBodyLimit(t *testing.T) {
}) })
addr, out := startProxy(t, app.URL, map[string]string{ addr, out := startProxy(t, app.URL, map[string]string{
requestMaxBytes: sizeLimitSetting, requestMaxBytes: sizeLimitSetting,
metricsToken: token,
}) })
var body io.Reader = bytes.NewReader(make([]byte, tc.size)) var body io.Reader = bytes.NewReader(make([]byte, tc.size))
@@ -66,6 +67,13 @@ func TestRequestBodyLimit(t *testing.T) {
tc.want) tc.want)
wantLine(t, out.requestLine(t), tc.want, tc.action) wantLine(t, out.requestLine(t), tc.want, tc.action)
hits := 0
if tc.action == requestlog.ActionTooLarge {
hits = 1
}
wantLimitHits(t, addr, requestMaxBytes, hits)
if tc.refusedBeforeApp && calls.Load() != 0 { if tc.refusedBeforeApp && calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load()) t.Errorf("the app was called %d times, want never", calls.Load())
} }
@@ -106,6 +114,7 @@ func TestResponseBodyLimit(t *testing.T) {
}) })
addr, out := startProxy(t, app.URL, map[string]string{ addr, out := startProxy(t, app.URL, map[string]string{
responseMaxBytes: sizeLimitSetting, responseMaxBytes: sizeLimitSetting,
metricsToken: token,
}) })
got := get(t, addr, "/download") got := get(t, addr, "/download")
@@ -123,6 +132,13 @@ func TestResponseBodyLimit(t *testing.T) {
if line.UpstreamStatus != http.StatusOK { if line.UpstreamStatus != http.StatusOK {
t.Errorf("log line has upstream_status %d", line.UpstreamStatus) t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
} }
hits := 0
if tc.action == requestlog.ActionTooLarge {
hits = 1
}
wantLimitHits(t, addr, responseMaxBytes, hits)
}) })
} }
} }
+484
View File
@@ -0,0 +1,484 @@
package proxy_test
import (
"io"
"net/http"
"net/http/httptest"
"net/netip"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
// token is the SWWAF_METRICS_TOKEN the tests set, and bearer how a
// request carries it.
token = "0123456789abcdef0123456789abcdef"
bearer = "Bearer " + token
)
func TestMetricsAreOffWhileTheTokenIsUnset(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, out := startProxy(t, app.URL, nil)
// An empty token does not match the unset one either.
for i, authorization := range []string{bearer, "Bearer ", ""} {
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
if authorization != "" {
req.Header.Set("Authorization", authorization)
}
wantStatus(t, do(t, req), http.StatusNotFound)
wantLine(t, out.requestLines(t, i+1)[i], http.StatusNotFound,
requestlog.ActionAdmin)
}
if calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load())
}
}
func TestMetricsNeedTheTokenAndOtherPathsAreNotFound(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
for i, tc := range []struct {
method, path, authorization string
status int
}{
{http.MethodGet, proxy.MetricsPath, "", http.StatusUnauthorized},
{
http.MethodGet, proxy.MetricsPath, "Bearer " + strings.ToUpper(token),
http.StatusUnauthorized,
},
{http.MethodGet, proxy.MetricsPath, "Basic " + token, http.StatusUnauthorized},
{http.MethodGet, proxy.MetricsPath, bearer, http.StatusOK},
{http.MethodGet, proxy.MetricsPath, "bearer " + token, http.StatusOK},
{http.MethodPost, proxy.MetricsPath, bearer, http.StatusNotFound},
{http.MethodGet, proxy.MetricsPath + "/", bearer, http.StatusNotFound},
{http.MethodGet, "/_smallwebwaf/bans", bearer, http.StatusNotFound},
{http.MethodPost, proxy.HealthPath, "", http.StatusNotFound},
} {
req := newRequest(t, tc.method, addr, tc.path, http.NoBody)
if tc.authorization != "" {
req.Header.Set("Authorization", tc.authorization)
}
got := do(t, req)
wantStatus(t, got, tc.status)
wantLine(t, out.requestLines(t, i+1)[i], tc.status, requestlog.ActionAdmin)
if tc.status == http.StatusUnauthorized &&
got.header.Get("WWW-Authenticate") != "Bearer" {
t.Errorf("%q was answered without WWW-Authenticate: Bearer",
tc.authorization)
}
if tc.status == http.StatusOK &&
!strings.Contains(string(got.body), "# TYPE smallwebwaf_requests_total counter") {
t.Errorf("the metrics are\n%s", got.body)
}
}
if calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load())
}
}
func TestMetricsAreAskedForThroughTheChecks(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
metricsToken: token,
rateLimitPerMinute: "1",
})
// Asking for the metrics counts toward the client's limit of one
// request a minute, so its next request breaks it, and bans it. A
// banned client is refused the metrics too.
s.scrape(client)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
s.requestWithHeader(client, proxy.MetricsPath, "Authorization: "+bearer,
http.StatusForbidden, requestlog.ActionBanned)
}
func TestMetricsCountTheTraffic(t *testing.T) {
t.Parallel()
arrived, release := make(chan struct{}), make(chan struct{})
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path == "/held" {
close(arrived)
<-release
}
_, _ = io.WriteString(w, "hello")
})
releaseApp := sync.OnceFunc(func() { close(release) })
t.Cleanup(releaseApp)
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
wantStatus(t, got, http.StatusOK)
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
out.requestLines(t, 2)
forward := `{action="forward",status_class="2xx"}`
notFound := `{action="admin",status_class="4xx"}`
// The request for the metrics is itself under way.
metrics := scrape(t, addr)
wantMetric(t, metrics, "smallwebwaf_requests_total"+forward, 1)
wantMetric(t, metrics, "smallwebwaf_requests_total"+notFound, 1)
wantMetric(t, metrics, "smallwebwaf_request_bytes_total"+forward, 3)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
float64(len("Not Found\n")))
wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
metric(t, metrics, "go_goroutines")
metric(t, metrics, "process_start_time_seconds")
// A request the app holds is under way until it ends.
httpClient := newClient(t)
held := newRequest(t, http.MethodGet, addr, "/held", http.NoBody)
ended := make(chan error, 1)
go func() {
res, err := httpClient.Do(held)
if err == nil {
err = readAnswer(res).err
}
ended <- err
}()
<-arrived
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
releaseApp()
err := <-ended
if err != nil {
t.Fatalf("held request: %v", err)
}
out.requestLines(t, 5)
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
}
func TestMetricsCountLimitsAndBans(t *testing.T) {
t.Parallel()
const (
scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
denied = "192.0.2.50" // in SWWAF_DENY_NETS
)
s, clk, _ := startWithClock(t, "", map[string]string{
metricsToken: token,
rateLimitPerMinute: "1",
rateLimitExemptNets: scraper,
denyNets: denied,
banResponse: "close",
limitBanDuration: "1h",
maxBanDuration: "2h",
})
// SWWAF_BAN_RESPONSE=close sends no status at all.
s.get(denied, 0, requestlog.ActionDenied)
// A first broken limit bans for an hour.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, 0, requestlog.ActionRateLimited)
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
clk.advance(time.Hour)
wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
// A limit broken again right after would ban for three hours, longer
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, 0, requestlog.ActionRateLimited)
metrics = s.scrape(scraper)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
// denied, client, and the scraper as of its earlier requests.
wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
}
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
s, clk, server := startWithClock(t, "", map[string]string{
metricsToken: token,
adminToken: adminSecret,
rateLimitExemptNets: scraper,
})
const admins = `smallwebwaf_bans_made_total{cause="admin"}`
wantMetric(t, s.scrape(scraper), admins, 0)
// As an admin's edit of bans.json that adds a ban is taken in.
server.Ledger.LoadEdit([]bans.Ban{{
Netblock: netip.MustParsePrefix(client + "/32"),
Start: clk.Now(),
}})
wantMetric(t, s.scrape(scraper), admins, 1)
// And a ban made through the endpoint.
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
wantMetric(t, s.scrape(scraper), admins, 2)
}
func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
t.Parallel()
const fromFR = "198.51.100.20"
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
_, _ = io.WriteString(w, "hello")
})
env := map[string]string{
trustedProxies: trustLocalhost,
metricsToken: token,
metricsTopN: "2",
deniedCountries: "kp",
}
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
// The answers are kept before the requests, so that none waits for
// GeoJS.
server.GeoJS.Load([]lookup.Answer{
keptAnswer(fromKP, "KP"), keptAnswer(fromDE, "DE"), keptAnswer(fromFR, "FR"),
})
lines := 0
send := func(from string, times, status int) {
t.Helper()
for range times {
req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc"))
req.Header.Set(forwardedFor, from)
wantStatus(t, do(t, req), status)
// Each is counted before the next is sent, so that the
// countries are ranked in the order sent.
lines++
out.requestLines(t, lines)
}
}
// With two countries of their own, the third is counted as other.
send(fromKP, 3, http.StatusForbidden)
send(fromDE, 2, http.StatusOK)
send(fromFR, 1, http.StatusOK)
metrics := scrape(t, addr)
lines++
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
float64(3*len("Forbidden\n")))
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
float64(len("hello")))
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
// Once FR is busier than DE, it takes DE's place: its series counts
// from then on, and DE's is gone.
send(fromFR, 3, http.StatusOK)
metrics = scrape(t, addr)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
}
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
t.Parallel()
geojs := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(geojs.Close)
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{
trustedProxies: trustLocalhost,
metricsToken: token,
deniedCountries: "kp",
})
// GeoJS fails, so the client counts as coming from an unknown country,
// which SWWAF_DENIED_COUNTRIES does not refuse.
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, fromDE)
wantStatus(t, do(t, req), http.StatusOK)
// The client stops waiting for GeoJS after a second, so GeoJS's
// failure can come after its request has ended.
deadline := time.Now().Add(waitLimit)
metrics := scrape(t, addr)
for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
time.Now().Before(deadline) {
time.Sleep(pollInterval)
metrics = scrape(t, addr)
}
wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
}
// keptAnswer returns GeoJS's answer that the client at addr is in
// country, given now.
func keptAnswer(addr, country string) lookup.Answer {
now := time.Now()
return lookup.Answer{
Client: netip.MustParsePrefix(addr + "/32"), Country: country,
Answered: now, Used: now,
}
}
// scrape asks smallwebwaf at addr for the metrics, with the token, and
// returns them.
func scrape(t *testing.T, addr string) string {
t.Helper()
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
req.Header.Set("Authorization", bearer)
got := do(t, req)
if got.status != http.StatusOK {
t.Fatalf("the metrics were answered %d", got.status)
}
return string(got.body)
}
// scrape asks for the metrics, with the token, from the client at from,
// and returns them.
func (s *sender) scrape(from string) string {
s.t.Helper()
_, metrics := s.requestWithHeader(from, proxy.MetricsPath, "Authorization: "+bearer,
http.StatusOK, requestlog.ActionAdmin)
return metrics
}
// metric returns the value of series in metrics, which are in the
// Prometheus text format. series is a name and its labels in the order of
// their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
// the test if there is no such series.
func metric(t *testing.T, metrics, series string) float64 {
t.Helper()
for line := range strings.Lines(metrics) {
value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ")
if !found {
continue
}
number, err := strconv.ParseFloat(value, 64)
if err != nil {
t.Fatalf("%s has the value %q", series, value)
}
return number
}
t.Fatalf("no series %s in the metrics:\n%s", series, metrics)
return 0
}
// wantMetric checks the value of series in metrics, as metric reads it.
func wantMetric(t *testing.T, metrics, series string, want float64) {
t.Helper()
got := metric(t, metrics, series)
if got != want {
t.Errorf("%s is %v, want %v", series, got, want)
}
}
// wantNoSeries checks that metrics have no series series.
func wantNoSeries(t *testing.T, metrics, series string) {
t.Helper()
if strings.Contains(metrics, "\n"+series+" ") {
t.Errorf("there is a series %s", series)
}
}
// wantLimitHits checks that the metrics of smallwebwaf at addr count hits
// requests that passed the size or time limit of the setting limit, with
// no series for it when hits is 0.
func wantLimitHits(t *testing.T, addr, limit string, hits int) {
t.Helper()
series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}`
metrics := scrape(t, addr)
if hits == 0 {
wantNoSeries(t, metrics, series)
return
}
wantMetric(t, metrics, series, float64(hits))
}
+203
View File
@@ -0,0 +1,203 @@
package proxy_test
import (
"bytes"
"io"
"net/http"
"net/netip"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// observe is the value of SWWAF_MODE for observe mode.
const observe = "observe"
func TestObserveModeForwardsWhatEnforceModeRefuses(t *testing.T) {
t.Parallel()
const (
denied = "192.0.2.50" // in SWWAF_DENY_NETS
banned = otherClient // under a ban read from bans.json
)
for _, tc := range []struct {
setting string // "" leaves SWWAF_MODE at its default
observe bool
}{
{"", false},
{"enforce", false},
{observe, true},
} {
t.Run(mode+"="+tc.setting, func(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
env := map[string]string{
rateLimitPerMinute: "1",
denyNets: denied,
deniedCountries: "kp",
}
if tc.setting != "" {
env[mode] = tc.setting
}
s, clk, server := startWithClock(t, geojsURL, env)
server.Ledger.Load([]bans.Ban{{
Netblock: netip.MustParsePrefix(banned + "/32"),
Start: clk.Now(),
Expires: clk.Now().Add(time.Hour),
}})
// fromDE's first request is within the limit of one a minute,
// and its second breaks it.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
for _, sent := range []struct{ from, refusal string }{
{denied, requestlog.ActionDenied},
{banned, requestlog.ActionBanned},
{fromKP, requestlog.ActionCountryDenied},
{fromDE, requestlog.ActionRateLimited},
} {
if !tc.observe {
line := s.get(sent.from, http.StatusForbidden, sent.refusal)
wantWouldAction(t, line, "")
continue
}
// Passed to the app, which answered it.
line := s.get(sent.from, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, sent.refusal)
if line.UpstreamStatus != http.StatusOK {
t.Errorf("log line has upstream_status %d, want 200",
line.UpstreamStatus)
}
}
})
}
}
func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{
mode: observe,
rateLimitPerMinute: "1",
})
kept := bans.Ban{
Netblock: netip.MustParsePrefix(otherClient + "/32"),
Start: clk.Now(),
Expires: clk.Now().Add(time.Hour),
Cause: bans.CauseAdmin,
}
server.Ledger.Load([]bans.Ban{kept})
// No ban sets client's counters back to zero, so each request after
// the first breaks the limit of one a minute.
s.get(client, http.StatusOK, requestlog.ActionForward)
for range 2 {
line := s.get(client, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionRateLimited)
if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit ||
line.BanExpires != "" {
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
"want minute, limit and none", line.LimitHit, line.Offence,
line.BanExpires)
}
}
// The ban read from bans.json refuses nothing, and so counts no
// refusal in its notes, but is kept.
line := s.get(otherClient, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionBanned)
if line.BanExpires != requestlog.FormatTime(kept.Expires) {
t.Errorf("log line has ban_expires %q, want %s", line.BanExpires,
requestlog.FormatTime(kept.Expires))
}
got := server.Ledger.Snapshot()
if len(got) != 1 || got[0] != kept {
t.Errorf("bans\n%+v\nwant only\n%+v", got, kept)
}
}
func TestObserveModeKeepsTheSizeLimitsAndTheToken(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
var calls atomic.Int32
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
calls.Add(1)
answerWithSize(w, 2*sizeLimit, true)
})
addr, out := startProxy(t, app.URL, map[string]string{
mode: observe,
trustedProxies: trustLocalhost,
denyNets: denied,
requestMaxBytes: sizeLimitSetting,
responseMaxBytes: sizeLimitSetting,
metricsToken: token,
})
// SWWAF_DENY_NETS would refuse each request; instead a size limit or
// the missing token does.
for i, tc := range []struct {
method, path string
body io.Reader
status int
action string
}{
{
http.MethodPost, "/upload", bytes.NewReader(make([]byte, 2*sizeLimit)),
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge,
},
{
http.MethodGet, "/download", http.NoBody,
http.StatusBadGateway, requestlog.ActionTooLarge,
},
{
http.MethodGet, proxy.MetricsPath, http.NoBody,
http.StatusUnauthorized, requestlog.ActionAdmin,
},
} {
req := newRequest(t, tc.method, addr, tc.path, tc.body)
req.Header.Set(forwardedFor, denied)
wantStatus(t, do(t, req), tc.status)
line := out.requestLines(t, i+1)[i]
wantLine(t, line, tc.status, tc.action)
wantWouldAction(t, line, requestlog.ActionDenied)
}
// The upload was refused before it reached the app.
if calls.Load() != 1 {
t.Errorf("the app was called %d times, want once", calls.Load())
}
}
// wantWouldAction checks the request log line's would_action, and that a
// line that should have none has no such field.
func wantWouldAction(t *testing.T, line logLine, want string) {
t.Helper()
got, present := line.fields["would_action"]
switch {
case want == "" && present:
t.Errorf("log line has would_action %v, want none", got)
case want != "" && got != want:
t.Errorf("log line has would_action %v, want %s", got, want)
}
}
+67 -41
View File
@@ -6,6 +6,8 @@ import (
"errors" "errors"
"io" "io"
"net/http" "net/http"
"os"
"reflect"
"slices" "slices"
"strings" "strings"
"sync/atomic" "sync/atomic"
@@ -14,6 +16,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -115,27 +118,34 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
} }
} }
// wantRequestFields checks the log line's fields about the request. // wantRequestFields checks the log line's fields about the request. Its
// time, its id and its timings are checked only for being there.
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()
want := requestlog.Line{ hostname, _ := os.Hostname()
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery, want := withTimings(line, requestlog.Line{
Protocol: "HTTP/1.1", Status: http.StatusTeapot, Type: requestType, Time: line.Time, Instance: hostname,
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent), ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent),
ResponseBytes: int64(received), UserAgent: "test-agent", ResponseBytes: int64(received), UserAgent: "test-agent",
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal, RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
DurationUpstreamTotal: line.DurationUpstreamTotal, ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
} UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
if line.Line != want { Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
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)
} }
_, err := time.Parse(time.RFC3339, line.Time) _, err := time.Parse(time.RFC3339, line.Time)
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 { if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
t.Errorf("log line has time %q and durations %v and %v", line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
line.Time, line.DurationTotal, line.DurationUpstreamTotal) t.Errorf("log line has time %q, request_id %q and durations %v and %v",
line.Time, line.RequestID, line.DurationTotal,
line.fields["duration_upstream_total"])
} }
} }
@@ -297,7 +307,7 @@ func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) {
} }
} }
func TestServerHasTheFixedLimits(t *testing.T) { func TestServerHasTheDefaultLimits(t *testing.T) {
t.Parallel() t.Parallel()
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false }) cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
@@ -319,37 +329,48 @@ func TestServerHasTheFixedLimits(t *testing.T) {
} }
} }
func TestRefusesHeadersOver32KiB(t *testing.T) { func TestRefusesHeadersOverTheLimit(t *testing.T) {
t.Parallel() t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, _ := startProxy(t, app.URL, nil)
// size counts every byte of the request: the request line, the
// headers and the blank line that ends them.
const (
start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: "
end = "\r\n\r\n"
)
for _, tc := range []struct { for _, tc := range []struct {
size int name string
want int env map[string]string
limit int
}{ }{
{size: 32 << 10, want: http.StatusOK}, {"by default", nil, 32 << 10},
{size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge}, {"as set", map[string]string{clientHeaderMaxBytes: "8K"}, 8 << 10},
} { } {
conn := dial(t, addr) t.Run(tc.name, func(t *testing.T) {
send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end) t.Parallel()
wantStatus(t, readResponse(t, conn), tc.want)
}
if calls.Load() != 1 { var calls atomic.Int32
t.Errorf("the app was called %d times, want once", calls.Load())
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, _ := startProxy(t, app.URL, tc.env)
// size counts every byte of the request: the request line,
// the headers and the blank line that ends them.
const (
start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: "
end = "\r\n\r\n"
)
for _, sent := range []struct{ size, want int }{
{tc.limit, http.StatusOK},
{tc.limit + 1, http.StatusRequestHeaderFieldsTooLarge},
} {
conn := dial(t, addr)
send(t, conn,
start+strings.Repeat("a", sent.size-len(start)-len(end))+end)
wantStatus(t, readResponse(t, conn), sent.want)
}
if calls.Load() != 1 {
t.Errorf("the app was called %d times, want once", calls.Load())
}
})
} }
} }
@@ -360,8 +381,13 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
addr, out := startProxy(t, "http://"+localhost+":1", nil) addr, out := startProxy(t, "http://"+localhost+":1", nil)
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway) wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
wantLine(t, out.requestLine(t), http.StatusBadGateway,
requestlog.ActionUpstreamError) line := out.requestLine(t)
wantLine(t, line, http.StatusBadGateway, requestlog.ActionUpstreamError)
// There never was a connection to the app, nor an answer from it.
wantTimings(t, line, "duration_total", "duration_checks",
"duration_upstream_total")
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool { logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
return line["type"] == "process" && line["msg"] == "request to the app failed" return line["type"] == "process" && line["msg"] == "request to the app failed"
+113 -47
View File
@@ -8,26 +8,17 @@ import (
"log" "log"
"log/slog" "log/slog"
"net/http" "net/http"
"strings"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) "sneak.berlin/go/smallwebwaf/internal/rules"
// The request line and headers a client may send, and how long a
// kept-open client connection may wait for its next request, are fixed
// rather than settings. The limit on the request line and headers is
// 32 KiB, but Go's server reads 4 KiB past its MaxHeaderBytes before it
// refuses, so MaxHeaderBytes is set 4 KiB lower. The idle time is longer
// than the 90 seconds after which traefik closes a connection it is not
// using, so traefik never sends a request on a connection smallwebwaf is
// closing.
const (
requestHeaderMaxBytes = 32<<10 - 4<<10
clientIdleTimeout = 120 * time.Second
) )
// How smallwebwaf keeps connections to the app open between requests. // How smallwebwaf keeps connections to the app open between requests.
@@ -36,10 +27,26 @@ const (
appIdleConnTimeout = 90 * time.Second appIdleConnTimeout = 90 * time.Second
) )
// adminPrefix starts the path of every request for smallwebwaf itself,
// which never reaches the app.
const adminPrefix = "/_smallwebwaf/"
// HealthPath is smallwebwaf's health endpoint, which the container's // HealthPath is smallwebwaf's health endpoint, which the container's
// health check asks. // health check asks.
const HealthPath = "/_smallwebwaf/healthz" const HealthPath = "/_smallwebwaf/healthz"
// MetricsPath is where the metrics are, for a request that carries
// SWWAF_METRICS_TOKEN.
const MetricsPath = "/_smallwebwaf/metrics"
// BansPath is where an admin lists and adds bans, and, followed by / and
// a client's address, lifts them, with SWWAF_ADMIN_TOKEN.
const BansPath = "/_smallwebwaf/bans"
// ClientsPath is where an admin asks what smallwebwaf knows of a client,
// by the client's address after it, with SWWAF_ADMIN_TOKEN.
const ClientsPath = "/_smallwebwaf/clients/"
// Params are what New needs. // Params are what New needs.
type Params struct { type Params struct {
Config *config.Config Config *config.Config
@@ -51,48 +58,87 @@ type Params struct {
// lookup.URL. GeoJS is asked only while a country list is set. // lookup.URL. GeoJS is asked only while a country list is set.
GeoJSURL string GeoJSURL string
// Now tells the time by which requests are counted for the rate // Now tells the time by which requests are counted for the rate
// limits and bans are made and run out, normally time.Now. // limits, bans are made and run out, and GeoJS's answers are kept,
// normally time.Now in UTC, the time the state files give.
Now func() time.Time Now func() time.Time
// Rules are the rule files' rules, which each request is checked
// against.
Rules *rules.Files
// Alerts receive the alert for each ban the proxy makes or makes
// permanent, and for GeoJS failing.
Alerts *alerts.Queue
}
// Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep, and the metrics.
type Server struct {
*http.Server
Ledger *bans.Ledger
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
Metrics *metrics.Metrics
} }
// New returns the server smallwebwaf runs: each request it reads passes // New returns the server smallwebwaf runs: each request it reads passes
// through the proxy. Go's server itself refuses headers over 32 KiB, with // through the proxy. Go's server itself refuses a request line and
// 431, closes a connection idle for 120 seconds, and applies // headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy // SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
// applies the timeouts and size limits from then on. // applies the timeouts and size limits from then on.
func New(params Params) *http.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)
h := &handler{
config: params.Config,
requestLog: params.RequestLog,
processLog: params.ProcessLog,
errorLog: errorLog,
transport: newTransport(),
now: params.Now,
metrics: m,
limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay,
}),
ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
MaxBanDuration: params.Config.MaxBanDuration,
AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans,
}),
geojs: lookup.New(lookup.Params{
URL: params.GeoJSURL,
Now: params.Now,
ProcessLog: params.ProcessLog,
Metrics: m,
Alerts: params.Alerts,
}),
rules: params.Rules,
alerts: params.Alerts,
}
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
m.AddRules(params.Rules)
return &http.Server{ return &Server{
Addr: params.Config.ListenAddr, Server: &http.Server{
Handler: &handler{ Addr: params.Config.ListenAddr,
config: params.Config, Handler: h,
requestLog: params.RequestLog, ReadHeaderTimeout: params.Config.ClientRequestTimeout,
processLog: params.ProcessLog, // Off is an IdleTimeout of 0, which Go's server replaces with
errorLog: errorLog, // ReadTimeout: no limit, as long as ReadTimeout stays unset.
transport: newTransport(), IdleTimeout: params.Config.ClientIdleTimeout,
now: params.Now, // Go's server reads 4 KiB past MaxHeaderBytes before it
limiter: ratelimit.New(ratelimit.Limits{ // refuses, so the limit a client meets is the setting.
PerMinute: params.Config.RateLimitPerMinute, MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
PerHour: params.Config.RateLimitPerHour, ErrorLog: errorLog,
PerDay: params.Config.RateLimitPerDay,
}),
ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
MaxBanDuration: params.Config.MaxBanDuration,
MaxBans: params.Config.MaxBans,
}),
geojs: lookup.New(lookup.Params{
URL: params.GeoJSURL,
Now: time.Now,
ProcessLog: params.ProcessLog,
}),
}, },
ReadHeaderTimeout: params.Config.ClientRequestTimeout, Ledger: h.ledger,
IdleTimeout: clientIdleTimeout, Limiter: h.limiter,
MaxHeaderBytes: requestHeaderMaxBytes, GeoJS: h.geojs,
ErrorLog: errorLog, Metrics: m,
} }
} }
@@ -105,9 +151,12 @@ type handler struct {
errorLog *log.Logger errorLog *log.Logger
transport http.RoundTripper transport http.RoundTripper
now func() time.Time now func() time.Time
metrics *metrics.Metrics
limiter *ratelimit.Limiter limiter *ratelimit.Limiter
ledger *bans.Ledger ledger *bans.Ledger
geojs *lookup.GeoJS geojs *lookup.GeoJS
rules *rules.Files
alerts *alerts.Queue
} }
// newTransport returns what carries requests to the app. It never goes // newTransport returns what carries requests to the app. It never goes
@@ -124,7 +173,8 @@ func newTransport() *http.Transport {
// ServeHTTP handles one request: it works out the client, runs the // ServeHTTP handles one request: it works out the client, runs the
// checks, passes the request to the app and the answer back within the // checks, passes the request to the app and the answer back within the
// limits, and writes the request's log line. // limits, or answers it itself if it is for smallwebwaf, and writes the
// request's log line.
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
rq := h.newRequest(w, r) rq := h.newRequest(w, r)
defer rq.finish() defer rq.finish()
@@ -133,17 +183,33 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// a health checker is never refused. It does not ask the app. // a health checker is never refused. It does not ask the app.
if r.Method == http.MethodGet && r.URL.Path == HealthPath { if r.Method == http.MethodGet && r.URL.Path == HealthPath {
rq.line.Action = requestlog.ActionAdmin rq.line.Action = requestlog.ActionAdmin
// Set here rather than left to Go's server, which would set it only
// after the log line has taken the response's headers.
rq.out.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = io.WriteString(rq.out, "ok\n") _, _ = io.WriteString(rq.out, "ok\n")
return return
} }
// Once the request has ended, before its log line is written.
defer rq.addToHistory()
refused := rq.check(r.Context()) refused := rq.check(r.Context())
rq.checked = time.Now()
if refused != nil { if refused != nil {
rq.answer(*refused) rq.answer(*refused)
return return
} }
// A request for smallwebwaf itself is answered where another would be
// passed to the app, so that it goes through every check first.
if strings.HasPrefix(r.URL.Path, adminPrefix) {
rq.answerAdmin()
return
}
rq.forward(r.Context()) rq.forward(r.Context())
} }
+64 -6
View File
@@ -14,9 +14,11 @@ import (
"testing" "testing"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
) )
const ( const (
@@ -35,6 +37,10 @@ const (
// localhost is where every test server listens, and so the address // localhost is where every test server listens, and so the address
// smallwebwaf sees each test's requests come from. // smallwebwaf sees each test's requests come from.
localhost = "127.0.0.1" localhost = "127.0.0.1"
// requestType is the type that marks a request log line.
requestType = "request"
// protocol is the protocol of every test's requests.
protocol = "HTTP/1.1"
) )
// shortTimeoutSetting is shortTimeout as a setting's value. // shortTimeoutSetting is shortTimeout as a setting's value.
@@ -45,9 +51,12 @@ var shortTimeoutSetting = shortTimeout.String()
// The settings the tests set. // The settings the tests set.
const ( const (
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT" clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT" clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT" upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT" upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
mode = "SWWAF_MODE"
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"
@@ -56,6 +65,7 @@ const (
denyNets = "SWWAF_DENY_NETS" denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
deniedCountries = "SWWAF_DENIED_COUNTRIES" deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES" allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE" banResponse = "SWWAF_BAN_RESPONSE"
@@ -64,6 +74,10 @@ const (
maxBanDuration = "SWWAF_MAX_BAN_DURATION" maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS" maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX" banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
rulesDir = "SWWAF_RULES_DIR"
) )
// output collects what smallwebwaf writes on stdout. // output collects what smallwebwaf writes on stdout.
@@ -80,6 +94,14 @@ func (o *output) Write(p []byte) (int, error) {
return o.buf.Write(p) return o.buf.Write(p)
} }
// text returns everything written so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// lines returns every line written so far, decoded. // lines returns every line written so far, decoded.
func (o *output) lines(t *testing.T) []map[string]any { func (o *output) lines(t *testing.T) []map[string]any {
t.Helper() t.Helper()
@@ -119,7 +141,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
var found []logLine var found []logLine
for _, fields := range o.lines(t) { for _, fields := range o.lines(t) {
if fields["type"] == "request" { if fields["type"] == requestType {
found = append(found, decodeLine(t, fields)) found = append(found, decodeLine(t, fields))
} }
} }
@@ -194,14 +216,29 @@ func startProxyWithGeoJS(
} }
// startProxyWithClock is startProxyWithGeoJS with requests counted and // startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well. // bans made by the time now tells, and returns the server as well. Unless
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
func startProxyWithClock( func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time, t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string, env map[string]string,
) (string, *output, *http.Server) { ) (string, *output, *proxy.Server) {
t.Helper() t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL} addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env)
return addr, out, server
}
// startProxyWithAlerts is startProxyWithClock, and returns the queue of
// the alerts the proxy raises as well, as the settings in env make it. No
// alert is sent from it: they wait in it, for the test to look at.
func startProxyWithAlerts(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *proxy.Server, *alerts.Queue) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()}
maps.Copy(settings, env) maps.Copy(settings, env)
cfg, err := config.FromEnvironment(func(name string) (string, bool) { cfg, err := config.FromEnvironment(func(name string) (string, bool) {
@@ -214,12 +251,33 @@ func startProxyWithClock(
} }
out := &output{} out := &output{}
processLog := requestlog.NewProcessLogger(out)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
})
if err != nil {
t.Fatalf("rule files: %v", err)
}
alertQueue := alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
server := proxy.New(proxy.Params{ server := proxy.New(proxy.Params{
Config: cfg, Config: cfg,
RequestLog: out, RequestLog: out,
ProcessLog: requestlog.NewProcessLogger(out), ProcessLog: processLog,
GeoJSURL: geojsURL, GeoJSURL: geojsURL,
Now: now, Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
}) })
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
@@ -235,7 +293,7 @@ func startProxyWithClock(
_ = server.Close() _ = server.Close()
}) })
return listener.Addr().String(), out, server return listener.Addr().String(), out, server, alertQueue
} }
// newClient returns an HTTP client that sends requests as they are made, // newClient returns an HTTP client that sends requests as they are made,
+88
View File
@@ -5,6 +5,8 @@ import (
"sync/atomic" "sync/atomic"
"testing" "testing"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -69,3 +71,89 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
t.Errorf("the app was called %d times, want 4", calls.Load()) t.Errorf("the app was called %d times, want 4", calls.Load())
} }
} }
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
s, _, server := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/,/favicon.ico",
denyNets: denied,
deniedCountries: "kp",
})
// The answers are kept before the requests, so that none waits for
// GeoJS.
server.GeoJS.Load([]lookup.Answer{
keptAnswer(client, "DE"), keptAnswer(fromKP, "KP"),
})
// With a limit of one request a minute, the requests for paths under a
// prefix are not counted, so client's first request for / is within
// the limit; and once client has reached it, they are not refused.
s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
if line.LimitHit != "" || line.Counts != (ratelimit.Counts{}) {
t.Errorf("log line has limit_hit %q and counts %+v, want neither",
line.LimitHit, line.Counts)
}
// A path outside every prefix is counted: /assets is not under
// /assets/, and breaks the limit.
s.request(client, "/assets", http.StatusForbidden, requestlog.ActionRateLimited)
// A ban, SWWAF_DENY_NETS and the country lists still refuse a path
// under a prefix.
s.request(client, "/assets/app.js", http.StatusForbidden, requestlog.ActionBanned)
s.request(denied, "/assets/app.js", http.StatusForbidden, requestlog.ActionDenied)
s.request(fromKP, "/assets/app.js",
http.StatusForbidden, requestlog.ActionCountryDenied)
}
func TestRateLimitCountsPathsThatAreNotExempt(t *testing.T) {
t.Parallel()
for _, sent := range []string{
// A prefix matches only at the start of the path.
"/static/assets/app.js",
// A prefix matches the path as sent: a router that matches the
// path as received does not take /%61ssets/x for a path under
// /assets/.
"/%61ssets/x",
// .. once percent-decoded: an app may act on these as /login, the
// last as a path under /sneak/app/ or as /assets/x.
"/assets/../login",
"/assets/%2e%2e/login",
"/assets/..%2Flogin",
"/assets/..;/login",
"/sneak/app/src/branch/main/..%2F..%2F..%2F..%2F..%2F..%2Fassets/x",
// Not under /assets/ as sent: Go's router takes /assets%2Fx for one
// path segment, not a path under /assets/.
"/assets%2Fx",
"/assets%2fx",
// Under /assets/ as sent, but holding an encoded slash, in either
// case, or a backslash: never exempt, whatever the prefix.
"/assets/x%2Fy",
"/assets/x%2fy",
`/assets/x\y`,
} {
t.Run(sent, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/",
})
// Counted, the second request breaks the limit of one request
// a minute.
s.request(client, sent, http.StatusOK, requestlog.ActionForward)
s.request(client, sent, http.StatusForbidden, requestlog.ActionRateLimited)
})
}
}
+268 -66
View File
@@ -7,11 +7,15 @@ import (
"net/http/httptrace" "net/http/httptrace"
"net/http/httputil" "net/http/httputil"
"net/netip" "net/netip"
"net/url"
"os" "os"
"slices"
"strings"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -21,11 +25,13 @@ const flushAfterEachWrite time.Duration = -1
// refusal is smallwebwaf refusing a request, or refusing to go on with it: // refusal is smallwebwaf refusing a request, or refusing to go on with it:
// the status the client is answered if the response has not started yet, // the status the client is answered if the response has not started yet,
// 0 to close the connection without an answer, and the action the log // 0 to close the connection without an answer, the action the log line
// line names. // names, and the setting whose size or time limit the request passed, if
// that is why.
type refusal struct { type refusal struct {
status int status int
action string action string
limit string
} }
// request is one request on its way through smallwebwaf, from the moment // request is one request on its way through smallwebwaf, from the moment
@@ -43,7 +49,9 @@ type request struct {
peer netip.Addr peer netip.Addr
peerTrusted bool peerTrusted bool
start time.Time start time.Time
// upstreamStart is when the request was handed to the app. // checked is when the checks were done, and upstreamStart when the
// request was handed to the app.
checked time.Time
upstreamStart time.Time upstreamStart time.Time
// cancel ends the request to the app. // cancel ends the request to the app.
cancel context.CancelFunc cancel context.CancelFunc
@@ -53,24 +61,34 @@ type request struct {
complete bool complete bool
// 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 from its own; once // own, and the transport starts and stops them, and notes the times
// timersStopped is set, none of them acts any more. // below, from its own; once timersStopped is set, none of the timeouts
// acts any more.
mu sync.Mutex mu sync.Mutex
timersStopped bool timersStopped bool
clientRequestTimer *time.Timer clientRequestTimer *time.Timer
upstreamRequestTimer *time.Timer upstreamRequestTimer *time.Timer
upstreamResponseTimer *time.Timer upstreamResponseTimer *time.Timer
// requestSent is when the app had been sent the whole request. // connected is when there was a connection to the app, requestSent
requestSent time.Time // when the app had been sent the whole request, and answerStarted
// when the first byte of its answer arrived.
connected time.Time
requestSent time.Time
answerStarted time.Time
} }
// newRequest starts handling r: it notes the time and works out the // newRequest starts handling r: it notes the time, counts the request as
// client. // under way, works out the client, and starts the log line with what is
// known of the request.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request { func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
h.metrics.RequestStarted()
start := time.Now() start := time.Now()
peer := peerAddress(r) peer := peerAddress(r)
trusted := h.config.TrustedProxies trusted := h.config.TrustedProxies
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted) peerTrusted := isInside(peer, trusted)
forwardedFor := r.Header.Values("X-Forwarded-For")
client := clientAddress(peer, forwardedFor, trusted)
rq := &request{ rq := &request{
h: h, h: h,
@@ -79,22 +97,37 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
out: &responseWriter{ResponseWriter: w}, out: &responseWriter{ResponseWriter: w},
client: client, client: client,
peer: peer, peer: peer,
peerTrusted: isInside(peer, trusted), peerTrusted: peerTrusted,
start: start, start: start,
line: requestlog.Line{ line: requestlog.Line{
Time: requestlog.FormatTime(start), Time: requestlog.FormatTime(start),
ClientIP: client.String(), Instance: h.config.InstanceName,
PeerIP: peer.String(), ClientIP: client.String(),
Method: r.Method, Method: r.Method,
Host: r.Host, Scheme: scheme(r, peerTrusted),
Path: r.URL.EscapedPath(), Host: r.Host,
Query: r.URL.RawQuery, Path: r.URL.EscapedPath(),
Protocol: r.Proto, Query: r.URL.RawQuery,
Referer: r.Referer(), Protocol: r.Proto,
UserAgent: r.UserAgent(), Referer: r.Referer(),
Action: requestlog.ActionForward, UserAgent: r.UserAgent(),
RequestID: requestID(r, peerTrusted),
PeerIP: peer.String(),
ForwardedFor: strings.Join(forwardedFor, ", "),
ClientGroup: clientGroup(client).String(),
ContentType: r.Header.Get("Content-Type"),
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
HasAuthorization: len(r.Header.Values("Authorization")) > 0,
HasCookie: len(r.Header.Values("Cookie")) > 0,
Action: requestlog.ActionForward,
}, },
} }
// A length of -1 is a body whose length was not announced.
if r.ContentLength > 0 {
rq.line.ContentLength = r.ContentLength
}
if r.Body != http.NoBody { if r.Body != http.NoBody {
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq} rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
} }
@@ -102,50 +135,124 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
return rq return rq
} }
// requestHeaders returns the headers of r that names lists, by name in
// lower case, each with its values joined by ", ". Authorization, Cookie
// and Set-Cookie are never among them, whatever names says.
func requestHeaders(r *http.Request, names []string) map[string]string {
headers := map[string]string{}
for _, name := range names {
switch name {
case "authorization", "cookie", "set-cookie":
continue
}
values := r.Header.Values(name)
if len(values) > 0 {
headers[name] = strings.Join(values, ", ")
}
}
return headers
}
// 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 its body is read or anything reaches the app. It
// returns nil to let the request through. A client in SWWAF_ALLOW_NETS // returns nil to let the request through. The checks of checkClient come
// skips every check but the size limit. For any other client, // first, answered with SWWAF_BAN_RESPONSE, or 403 for a block rule, and
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a // then the size limit, so that a request the rate limits count is counted
// client either refuses is not looked up, and then the country lists; a // even when it is refused for its size. In observe mode a request
// request any of them refuses is not counted for the rate limits. Then // checkClient refuses goes on to the size limit like any other. ctx is
// come the rate limits, unless the client is in // the request's own context.
// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted,
// one refused for its size too. Every refusal but the size limit's is
// answered with SWWAF_BAN_RESPONSE. ctx is the request's own context.
func (rq *request) check(ctx context.Context) *refusal { func (rq *request) check(ctx context.Context) *refusal {
cfg := rq.h.config action := rq.checkClient(ctx)
allowed := isInside(rq.client, cfg.AllowNets)
exempt := isInside(rq.client, cfg.RateLimitExemptNets)
now := rq.h.now()
if !allowed && isInside(rq.client, cfg.DenyNets) { switch {
return rq.banResponse(requestlog.ActionDenied) case action == "":
case rq.h.config.Observe:
// The log line names what enforce mode would have done.
rq.line.WouldAction = action
case action == requestlog.ActionRuleBlocked:
return &refusal{status: http.StatusForbidden, action: action}
default:
return rq.banResponse(action)
} }
if !allowed && rq.banned(now) { maxBytes := rq.h.config.RequestMaxBytes
return rq.banResponse(requestlog.ActionBanned)
}
if !allowed && rq.countryDenied(ctx) {
return rq.banResponse(requestlog.ActionCountryDenied)
}
if !allowed && !exempt && rq.limitBroken(now) {
return rq.banResponse(requestlog.ActionRateLimited)
}
maxBytes := cfg.RequestMaxBytes
if maxBytes > 0 && rq.in.ContentLength > maxBytes { if maxBytes > 0 && rq.in.ContentLength > maxBytes {
return &refusal{ return &refusal{
status: http.StatusRequestEntityTooLarge, status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge, action: requestlog.ActionTooLarge,
limit: "SWWAF_REQUEST_MAX_BYTES",
} }
} }
return nil return nil
} }
// checkClient runs the checks on the request's client, and returns the
// action of the first that refuses the request, or "" when none does. A
// client in SWWAF_ALLOW_NETS skips them. For any other client,
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
// client either refuses is not looked up, and then the country lists; 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 every other request is counted,
// and last the rule files. ctx is the request's own context.
func (rq *request) checkClient(ctx context.Context) string {
cfg := rq.h.config
if isInside(rq.client, cfg.AllowNets) {
return ""
}
now := rq.h.now()
if isInside(rq.client, cfg.DenyNets) {
return requestlog.ActionDenied
}
if rq.banned(now) {
return requestlog.ActionBanned
}
if rq.countryDenied(ctx) {
return requestlog.ActionCountryDenied
}
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if !exempt && rq.limitBroken(now) {
return requestlog.ActionRateLimited
}
return rq.checkRules(now)
}
// pathExempt reports whether the rate limits leave out a request for u
// because of SWWAF_RATE_LIMIT_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
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
// as one path segment, as Go's router does.
func pathExempt(u *url.URL, prefixes []string) bool {
decoded := u.Path
// EscapedPath is the path as the app receives it, not decoded.
sent := u.EscapedPath()
if strings.Contains(decoded, "..") || strings.Contains(decoded, `\`) ||
strings.Contains(strings.ToLower(sent), "%2f") {
return false
}
return slices.ContainsFunc(prefixes, func(prefix string) bool {
return strings.HasPrefix(sent, prefix)
})
}
// forward passes the request to the app and the app's answer back. ctx // forward passes the request to the app and the app's answer back. ctx
// is the request's own context. // is the request's own context.
func (rq *request) forward(ctx context.Context) { func (rq *request) forward(ctx context.Context) {
@@ -154,7 +261,9 @@ func (rq *request) forward(ctx context.Context) {
rq.cancel = cancel rq.cancel = cancel
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
WroteRequest: rq.wroteRequest, GotConn: rq.gotConn,
WroteRequest: rq.wroteRequest,
GotFirstResponseByte: rq.gotFirstResponseByte,
}) })
out := rq.in.WithContext(ctx) out := rq.in.WithContext(ctx)
@@ -177,7 +286,8 @@ func (rq *request) forward(ctx context.Context) {
} }
// rewrite makes the request the app receives: the client's request, // rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set. // unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set.
func (rq *request) rewrite(pr *httputil.ProxyRequest) { func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme pr.Out.URL.Scheme = upstream.Scheme
@@ -186,6 +296,7 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
// the query as the client sent it. // the query as the client sent it.
pr.Out.URL.RawQuery = pr.In.URL.RawQuery pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
} }
// modifyResponse looks at the app's answer before ReverseProxy passes it // modifyResponse looks at the app's answer before ReverseProxy passes it
@@ -199,13 +310,18 @@ func (rq *request) modifyResponse(res *http.Response) error {
// connection it takes over, not through rq.out. // connection it takes over, not through rq.out.
rq.stopTimers() rq.stopTimers()
rq.out.status = res.StatusCode rq.out.status = res.StatusCode
rq.line.Websocket = true
return nil return nil
} }
maxBytes := rq.h.config.ResponseMaxBytes maxBytes := rq.h.config.ResponseMaxBytes
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes { if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge}) rq.refuse(refusal{
status: http.StatusBadGateway,
action: requestlog.ActionTooLarge,
limit: "SWWAF_RESPONSE_MAX_BYTES",
})
return errResponseTooLarge return errResponseTooLarge
} }
@@ -271,13 +387,18 @@ func (rq *request) answer(r refusal) {
} }
// refuse records r, unless an earlier refusal was, and ends the request // refuse records r, unless an earlier refusal was, and ends the request
// to the app. // to the app, if one was made: smallwebwaf reads the body of a request
// it answers itself too.
func (rq *request) refuse(r refusal) { func (rq *request) refuse(r refusal) {
rq.refused.CompareAndSwap(nil, &r) rq.refused.CompareAndSwap(nil, &r)
rq.cancel()
if rq.cancel != nil {
rq.cancel()
}
} }
// finish ends the request's timeouts and writes its log line. // finish ends the request's timeouts, counts it in the metrics and writes
// its log line.
func (rq *request) finish() { func (rq *request) finish() {
rq.stopTimers() rq.stopTimers()
@@ -289,35 +410,90 @@ func (rq *request) finish() {
line := &rq.line line := &rq.line
line.Status = rq.out.status line.Status = rq.out.status
line.ResponseBytes = rq.out.bytes line.ResponseBytes = rq.out.bytes
header := rq.out.Header()
line.ResponseContentType = header.Get("Content-Type")
line.CacheControl = header.Get("Cache-Control")
line.Location = header.Get("Location")
if rq.body != nil { if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load() line.RequestBytes = rq.body.bytes.Load()
} }
// limit is the setting whose size or time limit the request passed.
var limit string
switch { switch {
case refused != nil: case refused != nil:
line.Action = refused.action line.Action = refused.action
limit = refused.limit
case errors.Is(rq.out.err, os.ErrDeadlineExceeded): case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to // The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
// take the response. // take the response.
line.Action = requestlog.ActionTimedOut line.Action = requestlog.ActionTimedOut
limit = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil): case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
line.Aborted = true line.Aborted = true
} }
now := time.Now() now := time.Now()
line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start)) duration := now.Sub(rq.start)
line.DurationTotal = requestlog.Milliseconds(duration)
line.DurationChecks = timing(rq.start, rq.checked)
var upstreamDuration time.Duration
if !rq.upstreamStart.IsZero() { if !rq.upstreamStart.IsZero() {
line.DurationUpstreamTotal = requestlog.Milliseconds(now.Sub(rq.upstreamStart)) upstreamDuration = now.Sub(rq.upstreamStart)
line.DurationUpstreamTotal = new(requestlog.Milliseconds(upstreamDuration))
rq.mu.Lock()
line.DurationUpstreamConnect = timing(rq.upstreamStart, rq.connected)
line.DurationUpstreamFirstByte = timing(rq.upstreamStart, rq.answerStarted)
rq.mu.Unlock()
} }
// Counted before the log line is written, so that the metrics count
// every request whose line is out.
rq.h.metrics.RequestEnded(line, limit, duration, upstreamDuration)
err := requestlog.Write(rq.h.requestLog, line) err := requestlog.Write(rq.h.requestLog, line)
if err != nil { if err != nil {
rq.h.processLog.Error("writing the request log failed", "error", err.Error()) rq.h.processLog.Error("writing the request log failed", "error", err.Error())
} }
} }
// timing is the time from start to end in milliseconds, for one of the
// log line's timings, or nil when end is zero: what it times never
// happened.
func timing(start, end time.Time) *float64 {
if end.IsZero() {
return nil
}
return new(requestlog.Milliseconds(end.Sub(start)))
}
// addToHistory adds the request, which has ended, to its client's
// history.
func (rq *request) addToHistory() {
var requestBytes int64
if rq.body != nil {
requestBytes = rq.body.bytes.Load()
}
forwarded := !rq.upstreamStart.IsZero()
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country,
Forwarded: forwarded,
Refused: !forwarded && rq.refused.Load() != nil,
Status: rq.out.status,
RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes,
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
})
}
// 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 {
@@ -350,21 +526,26 @@ func (rq *request) startRequestTimers() {
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 { if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
rq.clientRequestTimer = time.AfterFunc( rq.clientRequestTimer = time.AfterFunc(
time.Until(rq.clientRequestDeadline()), rq.requestTimedOut) time.Until(rq.clientRequestDeadline()), func() {
rq.requestTimedOut("SWWAF_CLIENT_REQUEST_TIMEOUT")
})
} }
timeout := rq.h.config.UpstreamRequestTimeout timeout := rq.h.config.UpstreamRequestTimeout
if timeout > 0 { if timeout > 0 {
rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut) rq.upstreamRequestTimer = time.AfterFunc(timeout, func() {
rq.requestTimedOut("SWWAF_UPSTREAM_REQUEST_TIMEOUT")
})
} }
} }
// requestTimedOut is called when a request timeout runs out while the // requestTimedOut is called when limit, SWWAF_CLIENT_REQUEST_TIMEOUT or
// request is still on its way to the app. The answer names the side // SWWAF_UPSTREAM_REQUEST_TIMEOUT, runs out while the request is still on
// smallwebwaf was waiting on at that moment: 408 when it was waiting for // its way to the app. The answer names the side smallwebwaf was waiting
// the client to send more of its body, 504 when it was waiting for the // on at that moment: 408 when it was waiting for the client to send more
// app to be reached or to take what it had. // of its body, 504 when it was waiting for the app to be reached or to
func (rq *request) requestTimedOut() { // take what it had.
func (rq *request) requestTimedOut(limit string) {
rq.mu.Lock() rq.mu.Lock()
defer rq.mu.Unlock() defer rq.mu.Unlock()
@@ -376,6 +557,7 @@ func (rq *request) requestTimedOut() {
rq.refuse(refusal{ rq.refuse(refusal{
status: http.StatusGatewayTimeout, status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut, action: requestlog.ActionTimedOut,
limit: limit,
}) })
return return
@@ -384,6 +566,7 @@ func (rq *request) requestTimedOut() {
rq.refuse(refusal{ rq.refuse(refusal{
status: http.StatusRequestTimeout, status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut, action: requestlog.ActionTimedOut,
limit: limit,
}) })
// The transport gives up on the app only once its Read of the // The transport gives up on the app only once its Read of the
// client's body returns, so that Read is ended now. The lock keeps // client's body returns, so that Read is ended now. The lock keeps
@@ -399,6 +582,24 @@ func (rq *request) bodyReceived() {
stopTimer(rq.clientRequestTimer) stopTimer(rq.clientRequestTimer)
} }
// gotConn is called once there is a connection to the app, a new one or
// one kept open from an earlier request.
func (rq *request) gotConn(httptrace.GotConnInfo) {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.connected = time.Now()
}
// gotFirstResponseByte is called once the first byte of the app's answer
// has arrived.
func (rq *request) gotFirstResponseByte() {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.answerStarted = time.Now()
}
// wroteRequest is called once the app has been sent the whole request: // wroteRequest is called once the app has been sent the whole request:
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts. // the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) { func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
@@ -433,6 +634,7 @@ func (rq *request) responseTimedOut() {
rq.refuse(refusal{ rq.refuse(refusal{
status: http.StatusGatewayTimeout, status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut, action: requestlog.ActionTimedOut,
limit: "SWWAF_UPSTREAM_RESPONSE_TIMEOUT",
}) })
} }
} }
+368
View File
@@ -0,0 +1,368 @@
package proxy_test
import (
"io"
"maps"
"math"
"net/http"
"reflect"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// requestIDHeader carries the request's id.
requestIDHeader = "X-Request-ID"
// instance is the SWWAF_INSTANCE_NAME a test sets.
instance = "fsn1app1/gitea"
// ipv6Client is a client on IPv6, and ipv6Group the netblock the rate
// limits count it as.
ipv6Client = "2001:db8::7"
ipv6Group = "2001:db8::/64"
)
func TestLogLineHasEachFieldWhereItApplies(t *testing.T) {
t.Parallel()
received := make(chan string, 2) // the request ids the app received
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path != "/full" {
w.WriteHeader(http.StatusNoContent)
return
}
w.Header().Set("Content-Type", "text/html")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Location", "/elsewhere")
w.WriteHeader(http.StatusFound)
_, _ = io.WriteString(w, "moved")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
rateLimitExemptNets: localhost,
instanceName: instance,
logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE",
})
// This request comes from ipv6Client through a trusted proxy, with a
// body and each header the log line looks at, and is answered with a
// redirect.
conn := dial(t, addr)
send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+
forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+
forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+
"Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+
"Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+
"Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+
"Set-Cookie: secret-set-cookie\r\n\r\na=b")
wantStatus(t, readResponse(t, conn), http.StatusFound)
// A request's log line can come after its answer: each is waited for
// before the next request, so that the lines are in order.
full := out.requestLines(t, 1)[0]
// This one comes from 127.0.0.1, which the rate limits do not count,
// with a body of 4 bytes whose length it does not announce, so that its
// request_bytes is not its content_length, and no header the log line
// looks at, and is answered with 204 and no header.
conn = dial(t, addr)
send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+
"Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusNoContent)
bare := out.requestLines(t, 2)[1]
wantFullLine(t, full)
wantBareLine(t, bare)
for _, line := range []logLine{full, bare} {
got := <-received
if got != line.RequestID {
t.Errorf("the app received request id %q, the log line has %q",
got, line.RequestID)
}
}
if strings.Contains(out.text(), "secret") {
t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s",
out.text())
}
}
// wantFullLine checks the log line of the request with every header the
// line looks at. Its timings are checked by TestTimingsAreInOrder.
func wantFullLine(t *testing.T, line logLine) {
t.Helper()
headers := map[string]string{"accept": "text/html", "x-custom": "one, two"}
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure,
Host: appHost, Path: "/full", Protocol: protocol,
Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5,
RequestID: "from-traefik", PeerIP: localhost,
ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group,
ContentType: "application/x-www-form-urlencoded", ContentLength: 3,
RequestHeaders: headers, HasAuthorization: true, HasCookie: true,
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
}
// wantBareLine checks the log line of the request with none of them, and
// that the fields that do not apply to it are left out.
func wantBareLine(t *testing.T, line logLine) {
t.Helper()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: localhost, Method: http.MethodPost, Scheme: plain,
Host: appHost, Path: "/bare", Protocol: protocol,
Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID,
PeerIP: localhost, ClientGroup: localhost + "/32",
UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward,
})
if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" {
t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want)
}
for _, name := range []string{
"forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"cache_control", "location", "counts",
} {
_, present := line.fields[name]
if present {
t.Errorf("log line has %s, which does not apply", name)
}
}
}
// withTimings returns want with the timings of line.
func withTimings(line logLine, want requestlog.Line) requestlog.Line {
want.DurationTotal = line.DurationTotal
want.DurationChecks = line.DurationChecks
want.DurationUpstreamConnect = line.DurationUpstreamConnect
want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte
want.DurationUpstreamTotal = line.DurationUpstreamTotal
return want
}
func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) {
t.Parallel()
const hasAuthorization, hasCookie = "has_authorization", "has_cookie"
for _, tc := range []struct{ header, field, other string }{
{"Authorization", hasAuthorization, hasCookie},
{"Cookie", hasCookie, hasAuthorization},
} {
t.Run("only "+tc.header, func(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out := startProxy(t, app.URL, nil)
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(tc.header, "secret")
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLine(t)
_, otherPresent := line.fields[tc.other]
if line.fields[tc.field] != true || otherPresent {
t.Errorf("log line has %s %v and %s %v, want true and none",
tc.field, line.fields[tc.field], tc.other, line.fields[tc.other])
}
})
}
}
func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) {
t.Parallel()
const sentID = "from-traefik"
sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}}
trusted := map[string]string{trustedProxies: trustLocalhost}
for _, tc := range []struct {
name string
env map[string]string
header http.Header
// wantID is the request id logged, "" for a new one.
wantID, wantScheme string
}{
{"a trusted proxy's are kept", trusted, sent, sentID, secure},
{"without them, the id is new and the scheme http", trusted, nil, "", plain},
{"another peer's are replaced", nil, sent, "", plain},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
received := make(chan string, 2)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
})
addr, out := startProxy(t, app.URL, tc.env)
// Two requests, so that two new ids can be told apart.
ids := make([]string, 0, 2)
for i := range 2 {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
maps.Copy(req.Header, tc.header)
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLines(t, i+1)[i]
ids = append(ids, line.RequestID)
got := <-received
if line.RequestID != got || line.Scheme != tc.wantScheme {
t.Errorf("log line has request_id %q and scheme %q, and the "+
"app received id %q; want the same id and scheme %q",
line.RequestID, line.Scheme, got, tc.wantScheme)
}
}
switch {
case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID):
t.Errorf("request ids %q, want %q", ids, tc.wantID)
case tc.wantID == "" && (slices.Contains(ids, sentID) ||
slices.Contains(ids, "") || ids[0] == ids[1]):
t.Errorf("request ids %q, want two new ones", ids)
}
})
}
}
func TestTimingsAreInOrder(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
// The pauses set the times apart; a hold-up of the test only
// lengthens them.
time.Sleep(time.Millisecond)
w.WriteHeader(http.StatusOK)
_ = http.NewResponseController(w).Flush()
time.Sleep(time.Millisecond)
_, _ = io.WriteString(w, "done")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
denyNets: denied,
})
// Each log line is waited for before the next request, so that the
// lines are in order.
wantStatus(t, get(t, addr, "/"), http.StatusOK)
forwarded := out.requestLines(t, 1)[0]
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, denied)
wantStatus(t, do(t, req), http.StatusForbidden)
refused := out.requestLines(t, 2)[1]
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
health := out.requestLines(t, 3)[2]
// A request passed to the app has every timing; one refused, none of
// the app's; the health check, which runs no check, only the total.
wantTimings(t, forwarded, "duration_total", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total")
wantTimings(t, refused, "duration_total", "duration_checks")
wantTimings(t, health, "duration_total")
if t.Failed() {
return
}
// In whole microseconds, as they are logged, so that the sum below is
// exact.
total := microseconds(forwarded.DurationTotal)
checks := microseconds(*forwarded.DurationChecks)
connect := microseconds(*forwarded.DurationUpstreamConnect)
firstByte := microseconds(*forwarded.DurationUpstreamFirstByte)
upstream := microseconds(*forwarded.DurationUpstreamTotal)
// The checks end before the request is handed to the app, and the
// connection comes before the answer, which the app ends after a
// pause.
if checks+upstream > total || connect >= firstByte || firstByte >= upstream {
t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+
"first byte %d, upstream total %d", total, checks, connect, firstByte,
upstream)
}
if *refused.DurationChecks > refused.DurationTotal {
t.Errorf("refused request's checks took %v of %v milliseconds",
*refused.DurationChecks, refused.DurationTotal)
}
}
// wantTimings checks that the timings named are the only ones line has.
func wantTimings(t *testing.T, line logLine, want ...string) {
t.Helper()
var got []string
for name := range line.fields {
if strings.HasPrefix(name, "duration_") {
got = append(got, name)
}
}
slices.Sort(got)
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want)
}
}
// microseconds is a timing in whole microseconds.
func microseconds(milliseconds float64) int64 {
return int64(math.Round(milliseconds * 1000))
}
func TestLogsAnUpgradedConnection(t *testing.T) {
t.Parallel()
app := startApp(t, echoAfterUpgrade)
addr, out := startProxy(t, app.URL, nil)
conn := dial(t, addr)
send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols)
_ = conn.Close()
line := out.requestLine(t)
if line.fields["websocket"] != true {
t.Errorf("log line has websocket %v, want true", line.fields["websocket"])
}
}
+233
View File
@@ -0,0 +1,233 @@
package proxy_test
import (
"net/http"
"net/netip"
"os"
"path/filepath"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// testRules are the rules most tests here load: a block rule for
// /blocked and a ban rule for /.env.
const testRules = `
blocked path block ^/blocked$
probe path ban ^/\.env$
`
func TestEachRuleAction(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, "noted path log ^/\n"+testRules),
banResponse: "429",
})
start := clk.Now()
// A log rule notes its match, and lets the request through.
line := s.get(client, http.StatusOK, requestlog.ActionForward)
wantRuleIDs(t, line, "noted")
// A block rule refuses with 403, whatever SWWAF_BAN_RESPONSE is, and
// bans no one.
line = s.request(client, "/blocked", http.StatusForbidden,
requestlog.ActionRuleBlocked)
wantRuleIDs(t, line, "noted", "blocked")
s.get(client, http.StatusOK, requestlog.ActionForward)
// A ban rule refuses with SWWAF_BAN_RESPONSE, and bans the client for
// seven days, the default.
line = s.request(client, "/.env", http.StatusTooManyRequests, requestlog.ActionBanned)
wantRuleIDs(t, line, "noted", "probe")
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: "matched the rule probe",
Notes: bans.Notes{
RuleID: "probe",
Target: "path",
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: "/.env",
Status: http.StatusTooManyRequests,
UserAgent: userAgent,
},
// The four requests up to and including the probe.
Requests: 4,
},
}
got := server.Ledger.Bans(netblock)
if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
// The next request is refused under the ban, without being checked
// against the rules, and makes the ban permanent.
clk.advance(time.Hour)
line = s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
wantRuleIDs(t, line)
if line.BanExpires != permanent {
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
}
clk.advance(365 * 24 * time.Hour)
s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
}
func TestNextClearSignOfAttackAfterABanBansPermanently(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
attackBanDuration: "1h",
})
// The first probe bans for SWWAF_ATTACK_BAN_DURATION.
line := s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != requestlog.FormatTime(clk.Now().Add(time.Hour)) {
t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires)
}
// Once that ban has run out without a request, the client is served,
// and its next probe bans it for good.
clk.advance(time.Hour)
s.get(client, http.StatusOK, requestlog.ActionForward)
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != permanent {
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
}
}
func TestRulesComeAfterTheOtherChecks(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, _, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
allowNets: allowed,
rateLimitExemptNets: exempt,
rateLimitPerMinute: "1",
})
// A client in SWWAF_ALLOW_NETS is not checked.
line := s.request(allowed, "/.env", http.StatusOK, requestlog.ActionForward)
wantRuleIDs(t, line)
// A probe over the rate limit breaks the limit before any rule sees
// it.
s.get(client, http.StatusOK, requestlog.ActionForward)
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionRateLimited)
wantRuleIDs(t, line)
limitBan := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(limitBan) != 1 || limitBan[0].Cause != bans.CauseLimit {
t.Errorf("bans %+v, want one for a broken limit", limitBan)
}
// A client the rate limits do not apply to is still checked.
s.get(exempt, http.StatusOK, requestlog.ActionForward)
s.request(exempt, "/.env", http.StatusForbidden, requestlog.ActionBanned)
}
func TestObserveModeLogsWhatTheRulesWouldDo(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
mode: observe,
})
line := s.request(client, "/blocked", http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionRuleBlocked)
wantRuleIDs(t, line, "blocked")
line = s.request(client, "/.env", http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionBanned)
wantRuleIDs(t, line, "probe")
if line.BanExpires != "" {
t.Errorf("log line has ban_expires %q, want none", line.BanExpires)
}
// No ban was made.
line = s.get(client, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, "")
if got := server.Ledger.Snapshot(); len(got) != 0 {
t.Errorf("bans %+v, want none", got)
}
}
func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
s, _, _ := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
metricsToken: token,
})
s.request(client, "/blocked", http.StatusForbidden, requestlog.ActionRuleBlocked)
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
wantMetric(t, metrics,
`smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
}
// writeRules writes content as a rule file into a new directory, and
// returns the directory.
func writeRules(t *testing.T, content string) string {
t.Helper()
dir := t.TempDir()
err := os.WriteFile(filepath.Join(dir, "test.rules"), []byte(content), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
return dir
}
// wantRuleIDs checks the request log line's rule_ids.
func wantRuleIDs(t *testing.T, line logLine, want ...string) {
t.Helper()
if !slices.Equal(line.RuleIDs, want) {
t.Errorf("log line has rule_ids %v, want %v", line.RuleIDs, want)
}
}
+39
View File
@@ -0,0 +1,39 @@
package proxy
import (
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// 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
// 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
// the client's netblock for a clear sign of attack, or in observe mode
// raises the alert for the ban it would have made.
func (rq *request) checkRules(now time.Time) string {
matched := rq.h.rules.Match(rq.in)
for _, rule := range matched {
rq.line.RuleIDs = append(rq.line.RuleIDs, rule.ID)
rq.h.metrics.RuleMatched(rule.ID, rule.Action)
}
if len(matched) == 0 {
return ""
}
// Only the last rule matched can refuse the request.
switch last := matched[len(matched)-1]; last.Action {
case rules.ActionBlock:
return requestlog.ActionRuleBlocked
case rules.ActionBan:
rq.banForAttack(now, last)
return requestlog.ActionBanned
default:
return ""
}
}
+49 -17
View File
@@ -28,7 +28,9 @@ func TestRequestTimeouts(t *testing.T) {
for _, tc := range []struct { for _, tc := range []struct {
name string name string
env map[string]string // limit is the setting set to shortTimeout, which runs out; long
// is one set to longTimeoutSetting, which does not, or "".
limit, long string
// appTakesNothing has the app never read, while the client sends // appTakesNothing has the app never read, while the client sends
// as fast as it can; otherwise the app reads, and the client // as fast as it can; otherwise the app reads, and the client
// stops sending halfway. // stops sending halfway.
@@ -36,30 +38,26 @@ func TestRequestTimeouts(t *testing.T) {
want int want int
}{ }{
{ {
name: "client request timeout, waiting on the client", name: "client request timeout, waiting on the client",
env: map[string]string{clientRequestTimeout: shortTimeoutSetting}, limit: clientRequestTimeout,
want: http.StatusRequestTimeout, want: http.StatusRequestTimeout,
}, },
{ {
name: "upstream request timeout, waiting on the client", name: "upstream request timeout, waiting on the client",
env: map[string]string{ limit: upstreamRequestTimeout,
upstreamRequestTimeout: shortTimeoutSetting, long: clientRequestTimeout,
clientRequestTimeout: longTimeoutSetting, want: http.StatusRequestTimeout,
},
want: http.StatusRequestTimeout,
}, },
{ {
name: "upstream request timeout, waiting on the app", name: "upstream request timeout, waiting on the app",
env: map[string]string{upstreamRequestTimeout: shortTimeoutSetting}, limit: upstreamRequestTimeout,
appTakesNothing: true, appTakesNothing: true,
want: http.StatusGatewayTimeout, want: http.StatusGatewayTimeout,
}, },
{ {
name: "client request timeout, waiting on the app", name: "client request timeout, waiting on the app",
env: map[string]string{ limit: clientRequestTimeout,
clientRequestTimeout: shortTimeoutSetting, long: upstreamRequestTimeout,
upstreamRequestTimeout: longTimeoutSetting,
},
appTakesNothing: true, appTakesNothing: true,
want: http.StatusGatewayTimeout, want: http.StatusGatewayTimeout,
}, },
@@ -84,7 +82,12 @@ func TestRequestTimeouts(t *testing.T) {
appURL, sendRequest = app.URL, sendPartOfBody appURL, sendRequest = app.URL, sendPartOfBody
} }
addr, out := startProxy(t, appURL, tc.env) env := map[string]string{tc.limit: shortTimeoutSetting, metricsToken: token}
if tc.long != "" {
env[tc.long] = longTimeoutSetting
}
addr, out := startProxy(t, appURL, env)
start := time.Now() start := time.Now()
got := readResponse(t, sendRequest(t, addr)) got := readResponse(t, sendRequest(t, addr))
wantTimedOut(t, start) wantTimedOut(t, start)
@@ -105,6 +108,7 @@ func TestRequestTimeouts(t *testing.T) {
wantStatus(t, got, want) wantStatus(t, got, want)
wantLine(t, out.requestLine(t), want, requestlog.ActionTimedOut) wantLine(t, out.requestLine(t), want, requestlog.ActionTimedOut)
wantLimitHits(t, addr, tc.limit, 1)
}) })
} }
} }
@@ -198,6 +202,7 @@ func TestAppTooSlowToAnswer(t *testing.T) {
}) })
addr, out := startProxy(t, app.URL, map[string]string{ addr, out := startProxy(t, app.URL, map[string]string{
upstreamResponseTimeout: shortTimeoutSetting, upstreamResponseTimeout: shortTimeoutSetting,
metricsToken: token,
}) })
start := time.Now() start := time.Now()
@@ -213,6 +218,8 @@ func TestAppTooSlowToAnswer(t *testing.T) {
t.Errorf("log line has upstream_status %v for an app that never answered", t.Errorf("log line has upstream_status %v for an app that never answered",
line.fields["upstream_status"]) line.fields["upstream_status"])
} }
wantLimitHits(t, addr, upstreamResponseTimeout, 1)
} }
func TestAppTooSlowToFinishItsAnswer(t *testing.T) { func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
@@ -261,6 +268,7 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
}) })
addr, out := startProxy(t, app.URL, map[string]string{ addr, out := startProxy(t, app.URL, map[string]string{
clientResponseTimeout: shortTimeoutSetting, clientResponseTimeout: shortTimeoutSetting,
metricsToken: token,
}) })
start := time.Now() start := time.Now()
@@ -272,4 +280,28 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
line := out.requestLine(t) line := out.requestLine(t)
wantTimedOut(t, start) wantTimedOut(t, start)
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut) wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
wantLimitHits(t, addr, clientResponseTimeout, 1)
}
func TestClosesAnIdleConnection(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, _ := startProxy(t, app.URL, map[string]string{
clientIdleTimeout: shortTimeoutSetting,
})
// The idle time starts once the answer is sent, so after start.
start := time.Now()
conn := dial(t, addr)
send(t, conn, "GET / HTTP/1.1\r\nHost: app\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusOK)
// The read deadline readResponse set still bounds this read.
_, err := conn.Read(make([]byte, 1))
if !errors.Is(err, io.EOF) {
t.Fatalf("read on the idle connection: %v, want it closed", err)
}
wantTimedOut(t, start)
} }
+119
View File
@@ -0,0 +1,119 @@
package ratelimit_test
import (
"net/netip"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
func TestHistoryKeepsEveryRequest(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for i, r := range []ratelimit.Request{
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5},
{Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response.
{Refused: true, Status: 0},
// Answered 404 at smallwebwaf's own endpoints: neither forwarded
// nor refused.
{Status: 404},
} {
limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
}
want := ratelimit.History{
FirstSeen: start,
LastSeen: start.Add(6 * time.Minute),
Country: "FR",
LookedUp: start.Add(3 * time.Minute),
Requests: 7,
Forwarded: 4,
Refused: 2,
RequestBytes: 15,
ResponseBytes: 122,
Responses: ratelimit.Responses{
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 2, Status5xx: 1,
},
Offences: ratelimit.Offences{Limit: 1},
}
got := historyOf(t, limiter, client)
if got != want {
t.Errorf("history\n%+v\nwant\n%+v", got, want)
}
}
func TestResetKeepsTheHistory(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
wantCount(t, limiter, client, start, "")
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
}
limiter.Reset(client)
if got := historyOf(t, limiter, client).Requests; got != limit {
t.Errorf("the history counts %d requests, want %d", got, limit)
}
}
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
for client, requests := range map[string]int{
"198.51.100.9/32": 2,
"198.51.100.10/32": 3,
"192.0.2.1/32": 5,
"2001:db8:5::/64": 7,
} {
for range requests {
limiter.AddToHistory(netip.MustParsePrefix(client), midnight(),
ratelimit.Request{})
}
}
for netblock, want := range map[string]int64{
"198.51.100.9/32": 2,
"198.51.100.0/24": 5,
"2001:db8:5::/64": 7,
"203.0.113.0/24": 0,
} {
got := limiter.Requests(netip.MustParsePrefix(netblock))
if got != want {
t.Errorf("%s has sent %d requests, want %d", netblock, got, want)
}
}
}
// historyOf returns client's history.
func historyOf(
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix,
) ratelimit.History {
t.Helper()
for _, c := range limiter.Snapshot() {
if c.Client == client {
return c.History
}
}
t.Fatalf("%s is not in the table", client)
return ratelimit.History{}
}
+301 -49
View File
@@ -1,11 +1,15 @@
// Package ratelimit counts each client's requests over a minute, an hour // Package ratelimit keeps the table of clients: each client's requests
// and a day, as the "Counting method" section of SPEC.md describes, and // counted over a minute, an hour and a day, as the "Counting method"
// tells when a request takes a client over a rate limit. The counts are // section of SPEC.md describes, which tell when a request takes the client
// kept in memory only, for at most 20,000 clients. // over a rate limit, and each client's history since it was first seen.
// At most 20,000 clients are kept, in memory, and written to clients.json
// and read from it by the state package.
package ratelimit package ratelimit
import ( import (
"net/http"
"net/netip" "net/netip"
"slices"
"sync" "sync"
"time" "time"
@@ -13,7 +17,8 @@ import (
) )
// maxClients is how many clients are kept. Past it, the least recently // maxClients is how many clients are kept. Past it, the least recently
// seen client is dropped, and starts afresh if it comes back. // seen client is dropped, with its history, and starts afresh if it comes
// back.
const maxClients = 20000 const maxClients = 20000
const day = 24 * time.Hour const day = 24 * time.Hour
@@ -26,20 +31,99 @@ type Limits struct {
PerDay int64 PerDay int64
} }
// Limiter counts each client's requests against the limits. It is safe // Limiter counts each client's requests against the limits, and keeps
// for concurrent use. // its history. It is safe for concurrent use.
type Limiter struct { type Limiter struct {
// windows are the minute, the hour and the day, in the order of
// Client.buckets.
windows [3]window windows [3]window
mu sync.Mutex mu sync.Mutex
// clients holds each client's buckets, one pair for each of windows, clients *simplelru.LRU[netip.Prefix, *Client]
// in the same order. }
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
// Client is a client in the table, as clients.json holds it: its buckets
// in each window, and its history.
type Client struct {
Client netip.Prefix `json:"client"`
Minute Buckets `json:"minute"`
Hour Buckets `json:"hour"`
Day Buckets `json:"day"`
History History `json:"history"`
}
// Buckets are a client's two buckets in one window: the requests in the
// bucket under way, which began at Start, and in the bucket before it.
type Buckets struct {
Start time.Time `json:"start"`
Current int64 `json:"current"`
Previous int64 `json:"previous"`
}
// History is what is known of a client since it was first seen.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type History struct {
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
// Country is the client's country as it was last looked up, and
// LookedUp when that was; both are empty while it never was.
Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to
// the app, Refused those refused before anything reached it, a 401 at
// smallwebwaf's own endpoints included, and neither the others
// smallwebwaf answered there.
Requests int64 `json:"requests"`
Forwarded int64 `json:"forwarded"`
Refused int64 `json:"refused"`
// RequestBytes and ResponseBytes are the body bytes of its requests
// and of the responses it was sent.
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
Responses Responses `json:"responses,omitzero"`
Offences Offences `json:"offences,omitzero"`
}
// Responses are the responses a client was sent, by status class;
// Status5xx counts every status from 500 up.
type Responses struct {
Status1xx int64 `json:"1xx,omitempty"`
Status2xx int64 `json:"2xx,omitempty"`
Status3xx int64 `json:"3xx,omitempty"`
Status4xx int64 `json:"4xx,omitempty"`
Status5xx int64 `json:"5xx,omitempty"`
}
// Offences are a client's offences, by kind.
type Offences struct {
// Limit is its requests that broke a rate limit.
Limit int64 `json:"limit"`
}
// Request is what a client's history keeps of one of its requests.
type Request struct {
// Country is the client's country, when the request looked it up.
Country string
// Forwarded is true for a request passed to the app, Refused for one
// refused before anything reached it, a 401 at smallwebwaf's own
// endpoints included. Both are false for any other request smallwebwaf
// answered there.
Forwarded bool
Refused bool
// Status is what the client was sent, 0 if nothing was.
Status int
// RequestBytes and ResponseBytes are the body bytes of the request
// and of its response.
RequestBytes int64
ResponseBytes int64
// BrokeLimit is true for a request that broke a rate limit.
BrokeLimit bool
} }
// New returns a Limiter for limits, with no client counted yet. // New returns a Limiter for limits, with no client counted yet.
func New(limits Limits) *Limiter { func New(limits Limits) *Limiter {
clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](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
} }
@@ -65,38 +149,197 @@ type Hit struct {
Requests float64 Requests float64
} }
// Counts are a client's requests in the minute, the hour and the day that
// end at a request, that request included.
type Counts struct {
Minute float64 `json:"minute"`
Hour float64 `json:"hour"`
Day float64 `json:"day"`
}
// 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. It reports whether the request takes the client over // not it is refused, and returns the client's requests in each window. It
// a limit, and the window whose limit it goes over, the shortest if it is // reports whether the request takes the client over a limit, and the
// over several. // window whose limit it goes over, the shortest if it is over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) { func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
counts, seen := l.clients.Get(client) var (
if !seen { requests [3]float64
counts = &[3]buckets{} hit Hit
l.clients.Add(client, counts) )
}
var hit Hit for i, b := range l.get(client).buckets() {
w := l.windows[i]
for i, w := range l.windows { requests[i] = b.add(now, w.length)
requests := counts[i].add(now, w.length) if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) { hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
} }
} }
return hit, hit.Window != "" counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
return counts, hit, hit.Window != ""
} }
// Reset sets client's counts in every window back to zero. // 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()
l.clients.Remove(client) c, seen := l.clients.Peek(client)
if seen {
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
}
}
// AddToHistory adds r, a request from client at now, to the client's
// history.
func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
l.mu.Lock()
defer l.mu.Unlock()
h := &l.get(client).History
if h.FirstSeen.IsZero() {
h.FirstSeen = now
}
h.LastSeen = now
if r.Country != "" {
h.Country = r.Country
h.LookedUp = now
}
h.Requests++
if r.Forwarded {
h.Forwarded++
}
if r.Refused {
h.Refused++
}
h.RequestBytes += r.RequestBytes
h.ResponseBytes += r.ResponseBytes
h.Responses.add(r.Status)
if r.BrokeLimit {
h.Offences.Limit++
}
}
// Requests returns how many requests the clients inside netblock have
// sent, as their histories count them.
func (l *Limiter) Requests(netblock netip.Prefix) int64 {
l.mu.Lock()
defer l.mu.Unlock()
// Most often the netblock is one client.
c, seen := l.clients.Peek(netblock)
if seen {
return c.History.Requests
}
var requests int64
for _, c := range l.clients.Values() {
if netblock.Overlaps(c.Client) {
requests += c.History.Requests
}
}
return requests
}
// Client returns client as the table holds it, and whether it does. It is
// not a request from client, and leaves when it was last seen unchanged.
func (l *Limiter) Client(client netip.Prefix) (Client, bool) {
l.mu.Lock()
defer l.mu.Unlock()
c, seen := l.clients.Peek(client)
if !seen {
return Client{}, false
}
return *c, true
}
// Len returns how many clients are in the table.
func (l *Limiter) Len() int {
l.mu.Lock()
defer l.mu.Unlock()
return l.clients.Len()
}
// Snapshot returns every client in the table, sorted by address, as
// clients.json lists them.
func (l *Limiter) Snapshot() []Client {
l.mu.Lock()
clients := make([]Client, 0, l.clients.Len())
for _, c := range l.clients.Values() {
clients = append(clients, *c)
}
l.mu.Unlock()
slices.SortFunc(clients, func(a, b Client) int {
return a.Client.Compare(b.Client)
})
return clients
}
// Load puts clients read from clients.json into the table, in place of
// the clients it holds, in the order they were last seen, so that the
// least recently seen is dropped first. Buckets whose time has passed at
// now are emptied.
func (l *Limiter) Load(clients []Client, now time.Time) {
clients = slices.Clone(clients)
slices.SortStableFunc(clients, func(a, b Client) int {
return a.History.LastSeen.Compare(b.History.LastSeen)
})
l.mu.Lock()
defer l.mu.Unlock()
l.clients.Purge()
for _, c := range clients {
for i, b := range c.buckets() {
// The window that ends at now covers neither bucket once it
// begins after the bucket under way has ended.
length := l.windows[i].length
if !now.Add(-length).Before(b.Start.Add(length)) {
*b = Buckets{}
}
}
l.clients.Add(c.Client, &c)
}
}
// get returns client's entry in the table, a new one if it has none, and
// makes it the most recently seen.
func (l *Limiter) get(client netip.Prefix) *Client {
c, seen := l.clients.Get(client)
if !seen {
c = &Client{Client: client}
l.clients.Add(client, c)
}
return c
}
// buckets returns c's buckets in the minute, the hour and the day.
func (c *Client) buckets() [3]*Buckets {
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
} }
// window is a length of time over which requests are counted, and the // window is a length of time over which requests are counted, and the
@@ -107,14 +350,6 @@ type window struct {
limit int64 limit int64
} }
// buckets are a client's two buckets in one window: the requests in the
// bucket under way, which began at start, and in the bucket before it.
type buckets struct {
start time.Time
current int64
previous int64
}
// add counts a request at now in a window of length, and returns the // add counts a request at now in a window of length, and returns the
// client's requests in the window that ends at now: those in the bucket // client's requests in the window that ends at now: those in the bucket
// under way, and those in the bucket before it weighted by how much of // under way, and those in the bucket before it weighted by how much of
@@ -125,27 +360,44 @@ type buckets struct {
// 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) float64 {
if now.Before(b.start.Add(-time.Second)) { if now.Before(b.Start.Add(-time.Second)) {
*b = buckets{} *b = Buckets{}
} }
start := now.Truncate(length) start := now.Truncate(length)
if start.After(b.start) { if start.After(b.Start) {
if start.Equal(b.start.Add(length)) { if start.Equal(b.Start.Add(length)) {
b.previous = b.current b.Previous = b.Current
} else { } else {
b.previous = 0 b.Previous = 0
} }
b.start = start b.Start = start
b.current = 0 b.Current = 0
} }
b.current++ b.Current++
elapsed := max(now.Sub(b.start), 0) elapsed := max(now.Sub(b.Start), 0)
covered := 1 - float64(elapsed)/float64(length) covered := 1 - float64(elapsed)/float64(length)
return float64(b.previous)*covered + float64(b.current) return float64(b.Previous)*covered + float64(b.Current)
}
// add counts a response with status in its class. A status of 0, for
// nothing sent, is not a response.
func (r *Responses) add(status int) {
switch {
case status >= http.StatusInternalServerError:
r.Status5xx++
case status >= http.StatusBadRequest:
r.Status4xx++
case status >= http.StatusMultipleChoices:
r.Status3xx++
case status >= http.StatusOK:
r.Status2xx++
case status >= http.StatusContinue:
r.Status1xx++
}
} }
+26 -3
View File
@@ -62,14 +62,14 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight() start := midnight()
for range limit { for range limit {
_, over := limiter.Count(client, start) _, _, over := limiter.Count(client, start)
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)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1} want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want { if !over || hit != want {
@@ -78,6 +78,29 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
} }
} }
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range 3 {
limiter.Count(client, start)
}
// 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
// requests, which count 2.25, and this one: 3.25. The day covers all
// four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want {
t.Errorf("counts %+v, want %+v", counts, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) { func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel() t.Parallel()
@@ -238,7 +261,7 @@ func wantCount(
) { ) {
t.Helper() t.Helper()
hit, _ := limiter.Count(client, now) _, hit, _ := limiter.Count(client, now)
if hit.Window != want { if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q", t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want) client, now.Format(time.RFC3339), hit.Window, want)
+122
View File
@@ -0,0 +1,122 @@
package ratelimit_test
import (
"net/netip"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
func TestSnapshotListsTheClientsByAddress(t *testing.T) {
t.Parallel()
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{})
for _, i := range []int{2, 3, 0, 1} {
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
}
snapshot := limiter.Snapshot()
got := make([]string, 0, len(snapshot))
for _, c := range snapshot {
got = append(got, c.Client.String())
}
if !slices.Equal(got, want) {
t.Errorf("snapshot %v, want %v", got, want)
}
counted := ratelimit.Buckets{Start: midnight(), Current: 1}
if snapshot[0].Minute != counted || snapshot[0].Day != counted {
t.Errorf("buckets %+v and %+v, want %+v", snapshot[0].Minute, snapshot[0].Day,
counted)
}
}
func TestLoadedCountsCarryOn(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
before := ratelimit.New(ratelimit.Limits{PerHour: limit})
for range limit {
wantCount(t, before, client, start, "")
}
// Loaded into a new limiter, as across a restart, the client has no
// fresh allowance.
later := start.Add(time.Minute)
after := ratelimit.New(ratelimit.Limits{PerHour: limit})
after.Load(before.Snapshot(), later)
wantCount(t, after, client, later, hour)
}
func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client {
t.Helper()
after := ratelimit.New(ratelimit.Limits{})
after.Load(limiter.Snapshot(), now)
return after.Snapshot()[0]
}
// Two minutes on, the window that ends then covers neither of the
// minute's buckets, which are emptied; the hour's and the day's stay,
// and so does the history.
got := loaded(start.Add(2 * time.Minute))
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
got.Day.Current != 1 || got.History.Requests != 1 {
t.Errorf("loaded two minutes on as %+v", got)
}
// A moment before, the window still covers some of the earlier one.
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
if got.Minute.Current != 1 {
t.Errorf("loaded just under two minutes on with minute buckets %+v",
got.Minute)
}
}
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
t.Parallel()
const maxClients = 20000
// 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
// one seen longest ago, and the one dropped.
clients := make([]ratelimit.Client, maxClients+1)
addr := netip.MustParseAddr("10.0.0.0")
for i := range clients {
clients[i].Client = netip.PrefixFrom(addr, addr.BitLen())
clients[i].History.LastSeen = midnight().Add(-time.Duration(i) * time.Second)
addr = addr.Next()
}
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Load(clients, midnight())
got := limiter.Snapshot()
if len(got) != maxClients || got[0].Client != clients[0].Client ||
got[maxClients-1].Client != clients[maxClients-1].Client {
t.Errorf("%d clients kept, from %s to %s; want %d, from %s to %s",
len(got), got[0].Client, got[len(got)-1].Client, maxClients,
clients[0].Client, clients[maxClients-1].Client)
}
}
+310
View File
@@ -0,0 +1,310 @@
// Package remotelog sends the lines smallwebwaf writes on stdout to the
// remote log endpoint, SWWAF_LOG_REMOTE_URL, as the "Request log" section
// of SPEC.md describes: each line as the message of an RFC 5424 syslog
// record, over UDP, TCP or TLS. Lines wait in a bounded buffer, so a slow
// or unreachable endpoint never holds up a request or stdout.
package remotelog
import (
"bytes"
"context"
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"log/slog"
"net"
"net/url"
"os"
"strconv"
"sync/atomic"
"syscall"
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The forms of SWWAF_LOG_REMOTE_URL, by its scheme.
const (
SchemeUDP = "syslog+udp"
SchemeTCP = "syslog+tcp"
SchemeTLS = "syslog+tls"
)
// A record's priority is the number of its facility times the number of
// severities there are, plus the number of its severity. Every record's
// severity is informational.
const (
severities = 8
informational = 6
)
const (
// dialTimeout bounds connecting to the endpoint, the TLS handshake
// included.
dialTimeout = 10 * time.Second
// After a failed attempt to connect, or a connection on which a record
// fails, the next attempt to connect is made a second later, and
// retryDelayFactor times as long after each further failure in a row,
// up to a minute. A connection that fails after it has stayed up for
// resetRetryDelayAfter ends the row.
firstRetryDelay = time.Second
retryDelayFactor = 2
maxRetryDelay = time.Minute
resetRetryDelayAfter = time.Minute
)
// Params are what New needs.
type Params struct {
// URL is the endpoint (SWWAF_LOG_REMOTE_URL): SchemeUDP, SchemeTCP or
// SchemeTLS, a host and a port.
URL *url.URL
// RootCAs are the certificates a SchemeTLS endpoint's certificate
// must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil for the host's.
RootCAs *x509.CertPool
// Buffer is the most lines held while they wait to be sent
// (SWWAF_LOG_REMOTE_BUFFER).
Buffer int
// Facility is the number of the records' syslog facility
// (SWWAF_LOG_REMOTE_FACILITY), and AppName their APP-NAME
// (SWWAF_LOG_REMOTE_APP_NAME).
Facility int
AppName string
}
// Sender sends lines to the endpoint. Write puts them in its buffer, and
// Run sends them from there.
type Sender struct {
url *url.URL
tlsConfig *tls.Config
// beforeTime and afterTime are the parts of every record's header
// before and after its time, as RFC 5424 lays the header out.
beforeTime string
afterTime string
// records is the buffer: each line's record, framed to be sent.
records chan []byte
sent atomic.Int64
dropped atomic.Int64
}
// New returns a Sender for the endpoint params.URL.
func New(params Params) *Sender {
hostname, err := os.Hostname()
if err != nil || hostname == "" {
hostname = "-" // RFC 5424's value for a field that has none
}
priority := params.Facility*severities + informational
return &Sender{
url: params.URL,
tlsConfig: &tls.Config{
RootCAs: params.RootCAs,
MinVersion: tls.VersionTLS12,
},
// The 1 is the version of the format. The process id, the message
// id and the structured data have no value.
beforeTime: "<" + strconv.Itoa(priority) + ">1 ",
afterTime: " " + hostname + " " + params.AppName + " - - - ",
records: make(chan []byte, params.Buffer),
}
}
// Write puts each line in p in the buffer, as the message of a record of
// its own, and never waits: when the buffer is full, the oldest record in
// it is dropped to make room. It is safe for concurrent use.
func (s *Sender) Write(p []byte) (int, error) {
at := requestlog.FormatTime(time.Now())
for line := range bytes.Lines(p) {
line = bytes.TrimSuffix(line, []byte("\n"))
if len(line) > 0 {
s.put(s.record(at, line))
}
}
return len(p), nil
}
// Sent is how many records have been sent.
func (s *Sender) Sent() int64 {
return s.sent.Load()
}
// Dropped is how many records were dropped: the oldest in a full buffer,
// and those whose sending failed.
func (s *Sender) Dropped() int64 {
return s.dropped.Load()
}
// Depth is how many records are in the buffer.
func (s *Sender) Depth() int {
return len(s.records)
}
// Run connects to the endpoint and sends each record as it comes into the
// buffer, until ctx is done. Then it sends the records still in the buffer,
// on the connection open at that time or, if there is none, on a new one,
// until none is left or one fails, and returns. How long it may take over
// that is for the caller to bound.
//
// A connection on which a record fails is closed and the record dropped.
// That failure, like a failed attempt to connect, is logged to processLog
// and followed by the next attempt after firstRetryDelay, retryDelayFactor
// times as long after each further failure in a row up to maxRetryDelay,
// and firstRetryDelay again after a connection that stayed up for
// resetRetryDelayAfter. Meanwhile the records wait in the buffer.
func (s *Sender) Run(ctx context.Context, processLog *slog.Logger) {
conn := s.send(ctx, processLog)
if conn == nil && len(s.records) > 0 {
conn, _ = s.dial(context.WithoutCancel(ctx))
}
if conn == nil {
return
}
defer func() {
_ = conn.Close()
}()
for {
select {
case record := <-s.records:
if s.write(conn, record) != nil {
return
}
default:
return
}
}
}
// record returns line as an RFC 5424 record made at the time at, framed
// for the endpoint: on its own over UDP, since each datagram holds one,
// and over TCP and TLS after its length in bytes and a space, the
// octet-counted framing of RFC 6587 and RFC 5425.
func (s *Sender) record(at string, line []byte) []byte {
record := make([]byte, 0, len(s.beforeTime)+len(at)+len(s.afterTime)+len(line))
record = append(record, s.beforeTime...)
record = append(record, at...)
record = append(record, s.afterTime...)
record = append(record, line...)
if s.url.Scheme == SchemeUDP {
return record
}
return append([]byte(strconv.Itoa(len(record))+" "), record...)
}
// put adds record to the buffer, first dropping the oldest record in it
// while it is full.
func (s *Sender) put(record []byte) {
for {
select {
case s.records <- record:
return
default:
}
select {
case <-s.records:
s.dropped.Add(1)
default:
}
}
}
// send connects to the endpoint and sends each record as it comes into
// the buffer, until ctx is done, and returns the connection then open, or
// nil.
func (s *Sender) send(ctx context.Context, processLog *slog.Logger) net.Conn {
delay := firstRetryDelay
for {
conn, err := s.dial(ctx)
if ctx.Err() != nil {
return conn
}
if err == nil {
connected := time.Now()
err = s.sendOn(ctx, conn)
if err == nil {
return conn
}
_ = conn.Close()
if time.Since(connected) >= resetRetryDelayAfter {
delay = firstRetryDelay
}
}
processLog.Warn("sending to SWWAF_LOG_REMOTE_URL failed",
"error", err.Error(), "connecting_again_in", delay.String())
select {
case <-time.After(delay):
case <-ctx.Done():
return nil
}
delay = min(retryDelayFactor*delay, maxRetryDelay)
}
}
// sendOn sends each record on conn as it comes into the buffer, until one
// fails, whose error it returns, or ctx is done.
func (s *Sender) sendOn(ctx context.Context, conn net.Conn) error {
for {
select {
case record := <-s.records:
err := s.write(conn, record)
if err != nil {
return err
}
case <-ctx.Done():
return nil
}
}
}
// write sends record on conn, and counts it as sent or, if that fails,
// as dropped. A record too long for one UDP datagram is dropped without
// an error, since the connection has not failed: a long request must not
// hold up the lines after it.
func (s *Sender) write(conn net.Conn, record []byte) error {
_, err := conn.Write(record)
if err != nil {
s.dropped.Add(1)
if errors.Is(err, syscall.EMSGSIZE) {
return nil
}
return fmt.Errorf("send a record: %w", err)
}
s.sent.Add(1)
return nil
}
// dial connects to the endpoint.
func (s *Sender) dial(ctx context.Context) (net.Conn, error) {
dialer := &net.Dialer{Timeout: dialTimeout}
switch s.url.Scheme {
case SchemeUDP:
return dialer.DialContext(ctx, "udp", s.url.Host)
case SchemeTLS:
tlsDialer := &tls.Dialer{NetDialer: dialer, Config: s.tlsConfig}
return tlsDialer.DialContext(ctx, "tcp", s.url.Host)
default:
return dialer.DialContext(ctx, "tcp", s.url.Host)
}
}
+640
View File
@@ -0,0 +1,640 @@
package remotelog_test
import (
"bufio"
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"fmt"
"io"
"log/slog"
"math/big"
"net"
"net/url"
"os"
"slices"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// The tests run in a synctest bubble, where the time package runs on a
// clock of the test's own, which starts at 2000-01-01T00:00:00Z: a wait
// lasts exactly as long as it should, however slowly the test process
// runs, and synctest.Wait returns once the sender has done all it can
// before time passes. The endpoint is a listener on the loopback address.
// A test reads from it only once the records are on their way, and checks
// the sender's counts first, since a goroutine of the bubble that waits on
// the network keeps that clock from moving on. For the same reason the
// endpoint that refuses connections, a tlsEndpoint, runs outside the
// bubble: a sender connecting over TLS waits on the endpoint's answer.
const (
// started is the time a record made as a test starts gives.
started = "2000-01-01T00:00:00.000Z"
appName = "fsn1app1/gitea"
// local0 is the number of the default facility, and local0Info the
// priority of its records.
local0 = 16
local0Info = "<134>"
// loopback is where the endpoints listen.
loopback = "127.0.0.1:0"
)
func TestRecordsOverUDPGoOnePerDatagram(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = endpoint.Close() })
sender, _, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
_, _ = sender.Write([]byte(`{"type":"request"}` + "\n" + `{"type":"process"}` + "\n"))
synctest.Wait()
wantCounts(t, sender, 2, 0, 0)
for _, line := range []string{`{"type":"request"}`, `{"type":"process"}`} {
datagram := make([]byte, 1024)
n, _, err := endpoint.ReadFrom(datagram)
if err != nil {
t.Fatalf("read: %v", err)
}
want := record(t, local0Info, appName, line)
if string(datagram[:n]) != want {
t.Errorf("datagram %q, want %q", datagram[:n], want)
}
}
})
}
func TestRecordsOverTCPAreOctetCountedWithTheirFacility(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
endpointParams := params(remotelog.SchemeTCP, endpoint.Addr())
endpointParams.Facility = 19 // local3
endpointParams.AppName = "gitea"
sender, _, _ := run(t, endpointParams)
_, _ = sender.Write([]byte("first\nsecond\n"))
synctest.Wait()
wantCounts(t, sender, 2, 0, 0)
frames := bufio.NewReader(accept(t, endpoint))
wantFrame(t, frames, record(t, "<158>", "gitea", "first"))
wantFrame(t, frames, record(t, "<158>", "gitea", "second"))
})
}
func TestStalledEndpointHoldsUpNoWriteAndOldestRecordsAreDropped(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
certificate, roots := testCertificate(t)
endpoint := listen(t)
endpointParams := params(remotelog.SchemeTLS, endpoint.Addr())
endpointParams.RootCAs = roots
endpointParams.Buffer = 3
sender, _, _ := run(t, endpointParams)
// The sender connects, and its TLS handshake waits for an answer
// the endpoint does not give yet.
conn := accept(t, endpoint)
var stdout bytes.Buffer
out := io.MultiWriter(&stdout, sender)
for i := range 5 {
_, _ = fmt.Fprintf(out, "line %d\n", i+1)
}
if stdout.String() != "line 1\nline 2\nline 3\nline 4\nline 5\n" {
t.Errorf("stdout has %q", stdout.String())
}
wantCounts(t, sender, 0, 2, 3)
// Once the endpoint answers, the three newest records are sent.
server := tls.Server(conn, &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
})
err := server.HandshakeContext(t.Context())
if err != nil {
t.Fatalf("handshake: %v", err)
}
synctest.Wait()
wantCounts(t, sender, 3, 2, 0)
frames := bufio.NewReader(server)
for _, line := range []string{"line 3", "line 4", "line 5"} {
wantFrame(t, frames, record(t, local0Info, appName, line))
}
})
}
func TestReconnectsWithBackoffAfterTheEndpointGoesAway(t *testing.T) {
t.Parallel()
certificate, roots := testCertificate(t)
endpoint := startTLSEndpoint(t, certificate)
synctest.Test(t, func(t *testing.T) {
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
endpointParams.RootCAs = roots
sender, logged, _ := run(t, endpointParams)
_, _ = sender.Write([]byte("one\n"))
synctest.Wait()
wantCounts(t, sender, 1, 0, 0)
conn := endpoint.next(t)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "one"))
// The endpoint goes away: it closes the connection, and refuses the
// next ones. The sender notices when a record fails, and tries to
// connect again a second later, then two seconds after that.
endpoint.refusing.Store(true)
_ = conn.Close()
writeUntilDropped(t, sender, 1)
sent := sender.Sent()
_, _ = sender.Write([]byte("two\n"))
time.Sleep(time.Second)
synctest.Wait()
endpoint.refusing.Store(false)
time.Sleep(2*time.Second - time.Nanosecond)
synctest.Wait()
wantCounts(t, sender, sent, 1, 1)
// The endpoint is back, and the record waiting is sent.
time.Sleep(time.Nanosecond)
synctest.Wait()
wantCounts(t, sender, sent+1, 1, 0)
conn = endpoint.next(t)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "two"))
wantRetries(t, logged, "1s", "2s")
})
}
func TestAConnectionClosedAtOnceIsMadeAgainAfterAGrowingDelay(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
// The endpoint closes each connection as soon as it takes it. The
// sender notices when a record fails, and connects again a second
// later, then two seconds after that, then four.
delays := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second}
for i, delay := range delays {
_ = accept(t, endpoint).Close()
writeUntilDropped(t, sender, int64(i+1))
wantConnectedAgainAfter(t, sender, delay)
}
wantRetries(t, logged, "1s", "2s", "4s")
})
}
func TestTheDelayStartsAgainAfterAConnectionThatStayedUpAMinute(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
_ = accept(t, endpoint).Close()
writeUntilDropped(t, sender, 1)
wantConnectedAgainAfter(t, sender, time.Second)
// A connection that fails just short of a minute after it was made
// leaves the delay growing.
conn := accept(t, endpoint)
time.Sleep(time.Minute - time.Nanosecond)
_ = conn.Close()
writeUntilDropped(t, sender, 2)
wantConnectedAgainAfter(t, sender, 2*time.Second)
// One that fails a minute after it was made starts it again from a
// second.
conn = accept(t, endpoint)
time.Sleep(time.Minute)
_ = conn.Close()
writeUntilDropped(t, sender, 3)
wantConnectedAgainAfter(t, sender, time.Second)
wantRetries(t, logged, "1s", "2s", "1s")
})
}
func TestALineTooLongForADatagramIsDroppedAlone(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = endpoint.Close() })
sender, logged, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
// With its header, the first line's record is longer than the 65507
// bytes a UDP datagram over IPv4 holds. It is dropped, nothing is
// logged, and the next line is sent at once.
_, _ = sender.Write([]byte(strings.Repeat("x", 65507) + "\nnext\n"))
synctest.Wait()
wantCounts(t, sender, 1, 1, 0)
wantRetries(t, logged)
datagram := make([]byte, 1024)
n, _, err := endpoint.ReadFrom(datagram)
if err != nil {
t.Fatalf("read: %v", err)
}
want := record(t, local0Info, appName, "next")
if string(datagram[:n]) != want {
t.Errorf("datagram %q, want %q", datagram[:n], want)
}
})
}
func TestRecordsWaitingAtTheStopAreSent(t *testing.T) {
t.Parallel()
certificate, roots := testCertificate(t)
endpoint := startTLSEndpoint(t, certificate)
synctest.Test(t, func(t *testing.T) {
// The endpoint refuses the sender's first connection: it fails to
// connect, and waits a second to try again.
endpoint.refusing.Store(true)
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
endpointParams.RootCAs = roots
sender, logged, stop := run(t, endpointParams)
synctest.Wait()
wantRetries(t, logged, "1s")
_, _ = sender.Write([]byte("one\ntwo\n"))
endpoint.refusing.Store(false)
// Stopped before that second is over, it connects to send them.
stop()
wantCounts(t, sender, 2, 0, 0)
frames := bufio.NewReader(endpoint.next(t))
wantFrame(t, frames, record(t, local0Info, appName, "one"))
wantFrame(t, frames, record(t, local0Info, appName, "two"))
})
}
// output collects what the sender logs.
type output struct {
mu sync.Mutex
buf bytes.Buffer
}
// Write adds lines the sender logs.
func (o *output) Write(p []byte) (int, error) {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.Write(p)
}
// text returns everything logged so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// params returns the settings of a Sender for the endpoint at addr, in
// the form scheme names: room for ten lines, the default facility, and
// appName.
func params(scheme string, addr net.Addr) remotelog.Params {
return remotelog.Params{
URL: &url.URL{Scheme: scheme, Host: addr.String()},
Buffer: 10,
Facility: local0,
AppName: appName,
}
}
// run runs a Sender with settings until the test ends or the function
// it returns is called, which waits for Run to return. It returns the
// Sender, and what it logs.
func run(t *testing.T, settings remotelog.Params) (*remotelog.Sender, *output, func()) {
t.Helper()
sender := remotelog.New(settings)
logged := &output{}
ctx, cancel := context.WithCancel(t.Context())
ran := make(chan struct{})
go func() {
sender.Run(ctx, slog.New(slog.NewJSONHandler(logged, nil)))
close(ran)
}()
stop := func() {
cancel()
<-ran
}
t.Cleanup(stop)
return sender, logged, stop
}
// listen returns a TCP listener on the loopback address, closed when the
// test ends.
func listen(t *testing.T) net.Listener {
t.Helper()
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = listener.Close() })
return listener
}
// accept returns the next connection to listener, closed when the test
// ends.
func accept(t *testing.T, listener net.Listener) net.Conn {
t.Helper()
conn, err := listener.Accept()
if err != nil {
t.Fatalf("accept: %v", err)
}
t.Cleanup(func() { _ = conn.Close() })
return conn
}
// tlsEndpoint is a syslog+tls endpoint on the loopback address, which a
// test starts outside its bubble. It keeps its listener until the test
// ends, and either takes each connection or refuses it.
type tlsEndpoint struct {
addr net.Addr
// refusing is set while the endpoint closes each connection before the
// TLS handshake, which fails the sender's attempt to connect.
refusing atomic.Bool
// conns are the connections it has taken, after the handshake.
conns chan net.Conn
}
// startTLSEndpoint starts a tlsEndpoint with certificate, which takes
// connections until it is told to refuse them.
func startTLSEndpoint(t *testing.T, certificate tls.Certificate) *tlsEndpoint {
t.Helper()
listener := listen(t)
endpoint := &tlsEndpoint{addr: listener.Addr(), conns: make(chan net.Conn, 10)}
config := &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
}
go func() {
for {
conn, err := listener.Accept()
if err != nil {
return
}
server := tls.Server(conn, config)
if endpoint.refusing.Load() || server.HandshakeContext(t.Context()) != nil {
_ = conn.Close()
continue
}
endpoint.conns <- server
}
}()
return endpoint
}
// next returns the next connection the endpoint has taken, closed when
// the test ends.
func (e *tlsEndpoint) next(t *testing.T) net.Conn {
t.Helper()
conn := <-e.conns
t.Cleanup(func() { _ = conn.Close() })
return conn
}
// record returns the record of line made as the test started, with the
// priority and the app name given.
func record(t *testing.T, priority, app, line string) string {
t.Helper()
hostname, err := os.Hostname()
if err != nil || hostname == "" {
hostname = "-"
}
return priority + "1 " + started + " " + hostname + " " + app + " - - - " + line
}
// wantFrame reads the next octet-counted frame from frames, and checks
// that it holds want.
func wantFrame(t *testing.T, frames *bufio.Reader, want string) {
t.Helper()
count, err := frames.ReadString(' ')
if err != nil {
t.Fatalf("read a frame's length: %v", err)
}
length, err := strconv.Atoi(strings.TrimSuffix(count, " "))
if err != nil {
t.Fatalf("frame starts %q, not with its length", count)
}
got := make([]byte, length)
_, err = io.ReadFull(frames, got)
if err != nil {
t.Fatalf("read a frame: %v", err)
}
if string(got) != want {
t.Errorf("frame %q, want %q", got, want)
}
}
// wantCounts checks the records sender has sent, dropped and holds in
// its buffer.
func wantCounts(
t *testing.T, sender *remotelog.Sender, sent, dropped int64, depth int,
) {
t.Helper()
if sender.Sent() != sent || sender.Dropped() != dropped || sender.Depth() != depth {
t.Fatalf("sent %d, dropped %d, %d in the buffer; want %d, %d and %d",
sender.Sent(), sender.Dropped(), sender.Depth(), sent, dropped, depth)
}
}
// writeUntilDropped writes a line at a time until the count of records
// sender has dropped reaches dropped. The records it sends on a
// connection the endpoint has closed are lost before one fails; how many
// depends on when the endpoint's host answers that the connection is
// gone.
func writeUntilDropped(t *testing.T, sender *remotelog.Sender, dropped int64) {
t.Helper()
for sender.Dropped() < dropped {
_, _ = sender.Write([]byte("lost\n"))
synctest.Wait()
}
}
// wantConnectedAgainAfter writes a line while the sender waits to connect
// again, and checks that it connects, and takes the line from the buffer,
// only once delay is over.
func wantConnectedAgainAfter(
t *testing.T, sender *remotelog.Sender, delay time.Duration,
) {
t.Helper()
_, _ = sender.Write([]byte("waiting\n"))
time.Sleep(delay - time.Nanosecond)
synctest.Wait()
if sender.Depth() != 1 {
t.Fatalf("connected again before %v", delay)
}
time.Sleep(time.Nanosecond)
synctest.Wait()
if sender.Depth() != 0 {
t.Fatalf("not connected again after %v", delay)
}
}
// wantRetries checks that the sender logged a failure, of an attempt to
// connect or of a connection, for each of delays, the time until the next
// attempt, in order, and logged nothing else.
func wantRetries(t *testing.T, logged *output, delays ...string) {
t.Helper()
var got []string
for line := range strings.Lines(logged.text()) {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil || fields["msg"] != "sending to SWWAF_LOG_REMOTE_URL failed" {
t.Fatalf("logged %q", line)
}
delay, _ := fields["connecting_again_in"].(string)
got = append(got, delay)
}
if !slices.Equal(got, delays) {
t.Errorf("logged failures to connect again in %v, want %v", got, delays)
}
}
// testCertificate returns a certificate for 127.0.0.1 that is its own
// CA, and a pool that holds it. It is valid on the bubble's clock, which
// starts at 2000-01-01T00:00:00Z.
func testCertificate(t *testing.T) (tls.Certificate, *x509.CertPool) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("generate a key: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "smallwebwaf test CA"},
NotBefore: time.Date(1999, 12, 31, 0, 0, 0, 0, time.UTC),
NotAfter: time.Date(2000, 1, 2, 0, 0, 0, 0, time.UTC),
IsCA: true,
BasicConstraintsValid: true,
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)},
}
der, err := x509.CreateCertificate(rand.Reader, template, template,
&key.PublicKey, key)
if err != nil {
t.Fatalf("create a certificate: %v", err)
}
certificate, err := x509.ParseCertificate(der)
if err != nil {
t.Fatalf("parse the certificate: %v", err)
}
roots := x509.NewCertPool()
roots.AddCert(certificate)
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots
}
+84 -25
View File
@@ -9,6 +9,8 @@ import (
"io" "io"
"log/slog" "log/slog"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
) )
// The action a request line names: what smallwebwaf did with the // The action a request line names: what smallwebwaf did with the
@@ -26,8 +28,12 @@ const (
// ActionRateLimited is a request refused because it took its client // ActionRateLimited is a request refused because it took its client
// 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.
ActionBanned = "banned" ActionBanned = "banned"
// ActionRuleBlocked is a request refused because it matched a block
// rule.
ActionRuleBlocked = "rule_blocked"
// ActionDenied is a request refused because its client is in // ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS. // SWWAF_DENY_NETS.
ActionDenied = "denied" ActionDenied = "denied"
@@ -45,28 +51,74 @@ const OffenceLimit = "limit"
// 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"
// Line is one request's line in the request log. The field names are // Line is one request's line in the request log. The field names, and
// those of the "Request log" section of SPEC.md. // their order, are those of the "Request log" section of SPEC.md. A field
// that may not apply to a request is left out of its line when it does
// not.
// //
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case //nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Line struct { type Line struct {
Type string `json:"type"` Type string `json:"type"`
Time string `json:"time"`
ClientIP string `json:"client_ip"` // The standard web log fields. Scheme is how the client reached
PeerIP string `json:"peer_ip"` // smallwebwaf, or the trusted proxy in front of it.
Country string `json:"country"` Time string `json:"time"`
Method string `json:"method"` Instance string `json:"instance"`
Host string `json:"host"` ClientIP string `json:"client_ip"`
Path string `json:"path"` Method string `json:"method"`
Query string `json:"query"` Scheme string `json:"scheme"`
Protocol string `json:"protocol"` Host string `json:"host"`
Status int `json:"status"` Path string `json:"path"`
UpstreamStatus int `json:"upstream_status,omitempty"` Query string `json:"query"`
RequestBytes int64 `json:"request_bytes"` Protocol string `json:"protocol"`
ResponseBytes int64 `json:"response_bytes"` Status int `json:"status"`
Referer string `json:"referer"` RequestBytes int64 `json:"request_bytes"`
UserAgent string `json:"user_agent"` ResponseBytes int64 `json:"response_bytes"`
Action string `json:"action"` Referer string `json:"referer"`
UserAgent string `json:"user_agent"`
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
// or a new one, and is sent on to the app. ForwardedFor is the
// X-Forwarded-For header as received. ClientGroup is the netblock the
// client is counted as.
RequestID string `json:"request_id"`
PeerIP string `json:"peer_ip"`
ForwardedFor string `json:"forwarded_for,omitempty"`
ClientGroup string `json:"client_group"`
Country string `json:"country"`
ContentType string `json:"content_type,omitempty"`
// ContentLength is the length of its body the request announced.
ContentLength int64 `json:"content_length,omitempty"`
// RequestHeaders are the headers SWWAF_LOG_REQUEST_HEADERS names that
// the request carried, by name in lower case.
RequestHeaders map[string]string `json:"request_headers,omitempty"`
HasAuthorization bool `json:"has_authorization,omitempty"`
HasCookie bool `json:"has_cookie,omitempty"`
// Websocket is true when the connection was upgraded, as for a
// WebSocket.
Websocket bool `json:"websocket,omitempty"`
// Response detail, from the headers of the answer: the app's, as
// passed on, or those of smallwebwaf's own. Aborted is true when the
// client went away early.
ResponseContentType string `json:"response_content_type,omitempty"`
UpstreamStatus int `json:"upstream_status,omitempty"`
CacheControl string `json:"cache_control,omitempty"`
Location string `json:"location,omitempty"`
Aborted bool `json:"aborted,omitempty"`
// The decision.
Action string `json:"action"`
// WouldAction is, in observe mode, the action enforce mode would have
// taken with a request it would have refused: ActionDenied,
// ActionBanned, ActionCountryDenied, ActionRateLimited or
// ActionRuleBlocked.
WouldAction string `json:"would_action,omitempty"`
// Counts are the client's requests as the rate limits counted them
// with this one, for a request they counted.
Counts ratelimit.Counts `json:"counts,omitzero"`
// RuleIDs are the ids of the rule file rules the request matched.
RuleIDs []string `json:"rule_ids,omitempty"`
// LimitHit is the window whose rate limit the request went over: // LimitHit is the window whose rate limit the request went over:
// minute, hour or day. // minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"` LimitHit string `json:"limit_hit,omitempty"`
@@ -75,11 +127,18 @@ type Line struct {
// BanExpires is when the ban the request made, or was refused under, // BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent". // ends: a time, or "permanent".
BanExpires string `json:"ban_expires,omitempty"` BanExpires string `json:"ban_expires,omitempty"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"` // The timings, in milliseconds. DurationChecks is the time until the
// DurationTotal and DurationUpstreamTotal are in milliseconds. // checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte
DurationTotal float64 `json:"duration_total"` // and DurationUpstreamTotal run from when the request was handed to the
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"` // 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
// a request that did not get that far.
DurationTotal float64 `json:"duration_total"`
DurationChecks *float64 `json:"duration_checks,omitempty"`
DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"`
DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"`
DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"`
} }
// Write writes line to w as one JSON line marked "type":"request". // Write writes line to w as one JSON line marked "type":"request".
+5 -1
View File
@@ -50,7 +50,11 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
} }
unset := []string{ unset := []string{
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted", "forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"upstream_status", "cache_control", "location", "aborted", "counts",
"limit_hit", "offence", "ban_expires", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total", "duration_upstream_total",
} }
for _, name := range unset { for _, name := range unset {
+483
View File
@@ -0,0 +1,483 @@
// Package rules reads the rule files: the plain text files in
// SWWAF_RULES_DIR, one rule to a line, that each request is checked
// against, as the "Rule files" section of SPEC.md describes. They are read
// at start, and again once the directory has had no change for a short
// time after one is edited, added or removed.
package rules
import (
"context"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net/http"
"os"
"path/filepath"
"regexp"
"slices"
"strings"
"sync/atomic"
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
// The actions a rule takes when it matches.
const (
// ActionLog notes the match in the request log, and does nothing else.
ActionLog = "log"
// ActionBlock refuses the request with 403.
ActionBlock = "block"
// ActionBan refuses the request and bans the client's netblock: the
// request is a clear sign of attack.
ActionBan = "ban"
)
// extension ends the name of every rule file.
const extension = ".rules"
// quietTime is how long SWWAF_RULES_DIR must go without a change before
// the rule files are read again, so that a file still being written, such
// as one saved in place, appended to or copied in with scp, is read only
// once whole.
const quietTime = 2 * time.Second
// headerTarget starts the target that is one request header,
// header:<Name>.
const headerTarget = "header:"
// escapeLength is the length of a percent escape, such as %2e.
const escapeLength = 3
var (
// ruleLine is a rule: four fields separated by spaces or tabs, of
// which the fourth, the regex, runs to the end of the line.
ruleLine = regexp.MustCompile(`^([^ \t]+)[ \t]+([^ \t]+)[ \t]+([^ \t]+)[ \t]+(.+)$`)
// idChars are the characters of a rule's id.
idChars = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
)
var (
errNotRule = errors.New(
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs")
errNotID = errors.New("is not an id of letters, digits, - and _")
errNotTarget = errors.New(
"is not path, query, uri, method, host, user_agent, referer or header:<Name>")
errNotHeaderName = errors.New(
"has a character after header: that no header name can have")
errHeaderTakenOut = errors.New(
"names a header that Go's HTTP server takes out of every request, " +
"so a rule never sees it")
errNotAction = errors.New("is not log, block or ban")
errNotRegex = errors.New("does not compile")
errUsedTwice = errors.New("is already the id of the rule at")
)
// Rule is one rule of a rule file.
type Rule struct {
// ID names the rule in the request log, the metrics and ban notes.
ID string
// Target is what the regex is matched against, such as path or
// header:Accept.
Target string
// Action is ActionLog, ActionBlock or ActionBan.
Action string
regex *regexp.Regexp
}
// Params are what Load needs.
type Params struct {
// Dir is the directory of the rule files (SWWAF_RULES_DIR).
Dir string
// Enabled is SWWAF_RULES_ENABLED: while it is false, no file is read
// and no rule loaded.
Enabled bool
// ProcessLog receives how many rules were read, and the error in a
// rule file edited while smallwebwaf runs.
ProcessLog *slog.Logger
// Alerts receive a file_error alert for that error.
Alerts *alerts.Queue
}
// Files are the rule files of a running smallwebwaf, and the rules read
// from them. They are safe for concurrent use.
type Files struct {
params Params
// rules are the rules loaded, in the order of their files' names, and
// then of their lines.
rules atomic.Pointer[[]Rule]
}
// Load reads the rules of every *.rules file in Dir, in the order of the
// files' names, unless Enabled is false. A Dir that cannot be read is an
// error, and so is a line that is not a rule, a header name with a
// character no header name can have, a rule for the Host or the
// Transfer-Encoding header, which Go's HTTP server takes out of every
// request, a regex that does not compile and an id used twice, each named
// with its file and line.
func Load(params Params) (*Files, error) {
f := &Files{params: params}
f.rules.Store(&[]Rule{})
if !params.Enabled {
return f, nil
}
rules, _, err := read(params.Dir)
if err != nil {
return nil, err
}
f.rules.Store(&rules)
f.logRead(len(rules))
return f, nil
}
// Match checks r against the rules, in order, and returns those it
// matches, up to the first whose action refuses it, block or ban, which
// is then the last one returned.
func (f *Files) Match(r *http.Request) []Rule {
var matched []Rule
for _, rule := range *f.rules.Load() {
if !rule.matches(r) {
continue
}
matched = append(matched, rule)
if rule.Action != ActionLog {
break
}
}
return matched
}
// Len returns how many rules are loaded.
func (f *Files) Len() int {
return len(*f.rules.Load())
}
// Watch watches Dir until ctx is done, and reads the rule files again
// once Dir has had no change for quietTime, after one is edited, added or
// removed, and after Watch starts watching. If they then hold an error,
// the rules stay as they were, the error is logged with its file and
// line, and the files are read again after the next change. If Dir cannot
// be watched, that is logged, and the rules stay as they were loaded.
// While Enabled is false, Watch returns at once.
func (f *Files) Watch(ctx context.Context) {
if !f.params.Enabled {
return
}
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the rule files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the rule files for edits",
"directory", f.params.Dir)
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
}
// readAfterChanges reads the rule files again once quietTime has passed
// without a change from events, until ctx is done, and logs the errors
// from errs. The wait starts at once, as if for a change, so that an edit
// saved after Load read the files, and before Dir was watched, is taken
// in too.
func (f *Files) readAfterChanges(
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
) {
quiet := time.NewTimer(quietTime)
defer quiet.Stop()
for {
select {
case <-ctx.Done():
return
case <-events:
quiet.Reset(quietTime)
case <-quiet.C:
f.readAgain()
case err := <-errs:
f.params.ProcessLog.Warn("watching the rule files failed",
"error", err.Error())
}
}
}
// readAgain reads the rule files again, in place of the rules loaded, or
// logs the error that keeps the rules as they were, and raises a
// file_error alert for it, for the file it is in.
func (f *Files) readAgain() {
rules, path, err := read(f.params.Dir)
if err != nil {
const kept = "a rule file has an error, and the rules stay as they were"
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: kept,
Detail: map[string]any{"file": path, "error": err.Error()},
})
f.params.ProcessLog.Error(kept, "error", err.Error())
return
}
f.rules.Store(&rules)
f.logRead(len(rules))
}
// logRead logs that the rule files were read, and how many rules they
// hold, which can be none.
func (f *Files) logRead(count int) {
f.params.ProcessLog.Info("read the rule files",
"directory", f.params.Dir, "rules", count)
}
// read returns the rules of every rule file in dir, in the order of the
// files' names, and then of their lines, or an error, with the path of the
// rule file it is in, or dir. A file whose name starts with a dot, such as
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's
// *.rules would not match it.
func read(dir string) ([]Rule, string, error) {
entries, err := os.ReadDir(dir)
if err != nil {
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
}
var rules []Rule
// places are where each id is, as "<file>, line <n>".
places := map[string]string{}
for _, entry := range entries {
name := entry.Name()
if entry.IsDir() || strings.HasPrefix(name, ".") || filepath.Ext(name) != extension {
continue
}
path := filepath.Join(dir, name)
rules, err = readFile(path, rules, places)
if err != nil {
return nil, path, err
}
}
return rules, "", nil
}
// readFile appends the rules of the rule file at path to rules. places
// are where each id read so far is, and gain those of the file.
func readFile(path string, rules []Rule, places map[string]string) ([]Rule, error) {
data, err := os.ReadFile(path) //nolint:gosec // a rule file, in SWWAF_RULES_DIR
if err != nil {
return nil, err
}
number := 0
for line := range strings.Lines(string(data)) {
number++
place := fmt.Sprintf("%s, line %d", path, number)
text := strings.TrimSuffix(strings.TrimSuffix(line, "\n"), "\r")
rule, isRule, err := parse(text)
if err != nil {
return nil, fmt.Errorf("%s: %w", place, err)
}
if !isRule {
continue
}
first, used := places[rule.ID]
if used {
return nil, fmt.Errorf("%s: the id %q %w %s", place, rule.ID, errUsedTwice, first)
}
places[rule.ID] = place
rules = append(rules, rule)
}
return rules, nil
}
// parse reads a line of a rule file. It returns false for a blank line
// and for a comment, a line that starts with #. Spaces and tabs at the
// end of the line are not part of its regex, so a line with only those
// after its action has no regex, and is not a rule.
func parse(line string) (Rule, bool, error) {
line = strings.Trim(line, " \t")
if line == "" || strings.HasPrefix(line, "#") {
return Rule{}, false, nil
}
fields := ruleLine.FindStringSubmatch(line)
if fields == nil {
return Rule{}, false, errNotRule
}
rule := Rule{ID: fields[1], Target: fields[2], Action: fields[3]}
headerName, isHeader := strings.CutPrefix(rule.Target, headerTarget)
switch {
case !idChars.MatchString(rule.ID):
return Rule{}, false, fmt.Errorf("the id %q %w", rule.ID, errNotID)
case !isTarget(rule.Target):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotTarget)
case isHeader && !config.IsHeaderName(headerName):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotHeaderName)
case strings.EqualFold(rule.Target, headerTarget+"Host"):
return Rule{}, false, fmt.Errorf(
"the target %q %w; the request's host is the target host",
rule.Target, errHeaderTakenOut)
case strings.EqualFold(rule.Target, headerTarget+"Transfer-Encoding"):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errHeaderTakenOut)
case !slices.Contains([]string{ActionLog, ActionBlock, ActionBan}, rule.Action):
return Rule{}, false, fmt.Errorf("the action %q %w", rule.Action, errNotAction)
}
regex, err := regexp.Compile(fields[4])
if err != nil {
return Rule{}, false, fmt.Errorf("the regex %w: %w", errNotRegex, err)
}
rule.regex = regex
return rule, true, nil
}
// isTarget reports whether target is one a rule may have.
func isTarget(target string) bool {
switch target {
case "path", "query", "uri", "method", "host", "user_agent", "referer":
return true
}
name, isHeader := strings.CutPrefix(target, headerTarget)
return isHeader && name != ""
}
// matches reports whether the rule's regex matches its target in r. For
// uri it is matched against the path and query as received, and against
// them once percent-decoded, so that an encoded probe cannot slip past.
func (rule Rule) matches(r *http.Request) bool {
if rule.Target == "uri" {
uri := pathAndQuery(r)
return rule.regex.MatchString(uri) || rule.regex.MatchString(decodeOnce(uri))
}
return rule.regex.MatchString(value(rule.Target, r))
}
// value returns what a rule with target, other than uri, is matched
// against in r: the path and the query as the client sent them, before
// any decoding or re-encoding, split at 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":
path, _, _ := strings.Cut(pathAndQuery(r), "?")
return path
case "query":
_, query, _ := strings.Cut(pathAndQuery(r), "?")
return query
case "method":
return r.Method
case "host":
return r.Host
case "user_agent":
return header(r, "User-Agent")
case "referer":
return header(r, "Referer")
default:
return header(r, strings.TrimPrefix(target, headerTarget))
}
}
// pathAndQuery returns the target of r's request line, r.RequestURI, as
// the client sent it, less any scheme and host: a target with a scheme
// gives what follows the scheme and its :, and the host when // follows.
// So http://host/path, as a client sends it to a proxy, gives /path, and
// so does http:/path, which Go reads as a target with a scheme and no
// host. r.URL is not used: when the path holds a character it escapes,
// such as \ or a non-ASCII byte, it decodes the whole path and escapes it
// again, so that \ becomes %5C and %2e a dot.
func pathAndQuery(r *http.Request) string {
if !r.URL.IsAbs() {
return r.RequestURI
}
_, afterScheme, _ := strings.Cut(r.RequestURI, ":")
hostAndRest, hasHost := strings.CutPrefix(afterScheme, "//")
if !hasHost {
return afterScheme
}
start := strings.IndexAny(hostAndRest, "/?")
if start < 0 {
return ""
}
return hostAndRest[start:]
}
// header returns the values of r's header name joined by ", ", or "" if
// r has no such header.
func header(r *http.Request, name string) string {
return strings.Join(r.Header.Values(name), ", ")
}
// decodeOnce returns s with each percent escape, such as %2e, replaced by
// the byte it stands for. A % that is not followed by two hex digits is
// left as it is, so that a malformed escape cannot keep the rest of s
// from being decoded.
func decodeOnce(s string) string {
var decoded strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '%' && i+escapeLength <= len(s) {
b, err := hex.DecodeString(s[i+1 : i+escapeLength])
if err == nil {
decoded.Write(b)
i += escapeLength - 1
continue
}
}
decoded.WriteByte(s[i])
}
return decoded.String()
}
+687
View File
@@ -0,0 +1,687 @@
package rules_test
import (
"context"
"encoding/json"
"log/slog"
"maps"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"slices"
"strconv"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
const (
// What the process log says once Watch watches the directory, after
// each reading of the rule files, and for one that has an error.
watching = "watching the rule files for edits"
read = "read the rule files"
hasError = "a rule file has an error, and the rules stay as they were"
// maxLogLines is how many lines of the process log wait for a test to
// read them.
maxLogLines = 64
// browser is the user agent of an ordinary visitor.
browser = "Mozilla/5.0 (X11; Linux x86_64; rv:140.0) Gecko/20100101 Firefox/140.0"
// testFile is the rule file of a test that needs only one, and
// firstFile the first of a test's rule files.
testFile = "test.rules"
firstFile = "00-a.rules"
// userAgent is the header that carries the user agent.
userAgent = "User-Agent"
)
func TestEachTargetMatchesWhatItNames(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
rule string // its target, action and regex
uri string // the request's path and query
header http.Header
want bool
}{
{"path as received", `path log ^/%2eenv$`, "/%2eenv", nil, true},
{"path not decoded", `path log ^/\.env$`, "/%2eenv", nil, false},
{"path without the query", `path log ^/a$`, "/a?b=c", nil, true},
{"query as received", `query log ^b=%2e$`, "/a?b=%2e", nil, true},
{"uri as received", `uri log ^/a\?b=%2e$`, "/a?b=%2e", nil, true},
{"uri decoded", `uri log (\.\./){2}`, "/a?f=%2e%2e%2f%2e%2e%2f", nil, true},
{
"uri decoded past malformed escapes", `uri log (\.\./){2}&h=%$`,
"/a?g=%zz&f=%2e%2e%2f%2e%2e%2f&h=%", nil, true,
},
{"uri decoded only once", `uri log ^/a\.b$`, "/a%252eb", nil, false},
{"method", `method log ^PUT$`, "/", nil, true},
{"host", `host log ^app\.example$`, "/", nil, true},
{
"user_agent", `user_agent log ^sqlmap/`, "/",
http.Header{userAgent: {"sqlmap/1.8"}}, true,
},
{
"user_agent sent twice", `user_agent log ^curl/8, sqlmap/`, "/",
http.Header{userAgent: {"curl/8", "sqlmap/1.8"}}, true,
},
{"user_agent missing", `user_agent log ^$`, "/", nil, true},
{
"referer", `referer log ^https://spam\.example/`, "/",
http.Header{"Referer": {"https://spam.example/buy"}}, true,
},
{
"a header sent twice", `header:x-api-version log ^2, 3$`, "/",
http.Header{"X-Api-Version": {"2", "3"}}, true,
},
{"a header missing", `header:X-Api-Version log ^$`, "/", nil, true},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: "a-rule " + tc.rule + "\n"})
// Every request is a PUT, which the method rule looks for.
r := httptest.NewRequestWithContext(t.Context(), http.MethodPut,
"http://app.example"+tc.uri, nil)
maps.Copy(r.Header, tc.header)
got := len(files.Match(r)) == 1
if got != tc.want {
t.Errorf("%s matches %s: %t, want %t", tc.rule, tc.uri, got, tc.want)
}
})
}
}
func TestPathMatchedAsTheClientSentIt(t *testing.T) {
t.Parallel()
// Each path holds a character Go's URL type would escape again, \ or
// a non-ASCII byte, and each rule is written for the path as sent.
for _, tc := range []struct {
rule string // its target, action and regex
sent string // the path and query the client sent
}{
{`path log ^/\.\.\\\.\.\\windows\\win\.ini$`, `/..\..\windows\win.ini`},
{`path log ^/%2e%2e\\%2e%2e\\windows\\win\.ini$`, `/%2e%2e\%2e%2e\windows\win.ini`},
{`path log ^/café$`, "/café?x=1"},
{`uri log ^/%2e%2e\\%2e%2e\\boot\.ini\?x=1$`, `/%2e%2e\%2e%2e\boot.ini?x=1`},
} {
files := load(t, ruleFiles{testFile: "as-sent " + tc.rule + "\n"})
// The target in origin form, as traefik sends it, in absolute form,
// as a client sends it to a proxy, and with a scheme but no host,
// which Go reads as absolute form with no host, sending the app
// the path.
for _, target := range []string{
tc.sent, "http://app.example" + tc.sent, "http:" + tc.sent, "foo:" + tc.sent,
} {
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet, target, nil)
wantMatched(t, files, r, "as-sent")
}
}
}
func TestMatchingStopsAtTheFirstRuleThatRefuses(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: `
every-path path log ^/
no-path path log ^$
first-refusal path block ^/probe
later-ban path ban ^/probe
after path log ^/
`})
// Every log rule that matches is noted, and the block rule ends the
// matching.
wantMatched(t, files, get(t, "/probe"), "every-path", "first-refusal")
wantMatched(t, files, get(t, "/page"), "every-path", "after")
// A ban rule ends it too.
files = load(t, ruleFiles{testFile: "ban path ban ^/\nlater path block ^/\n"})
wantMatched(t, files, get(t, "/"), "ban")
}
func TestSpacesAndTabsEndingALineAreNotPartOfItsRegex(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: "env-file path block ^/\\.env$ \t \n"})
wantMatched(t, files, get(t, "/.env"), "env-file")
}
func TestFilesReadInNameOrderThenLineOrder(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{
"50-b.rules": "b1 path log ^/\n\n# a comment\n # an indented one\nb2 path log ^/\n",
firstFile: "a1 path log ^/\r\n",
// None is a rule file.
"notes.txt": "notes, not rules\n",
"10-c.rules.bak": "an old copy\n",
"20-d.rules/keep": "a file in a directory\n",
})
wantMatched(t, files, get(t, "/"), "a1", "b1", "b2")
if files.Len() != 3 {
t.Errorf("%d rules loaded, want 3", files.Len())
}
}
func TestFileWhoseNameStartsWithADotIsNotARuleFile(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "probe path block ^/probe\n"})
// The lock file Emacs makes beside a file while it is edited: a link to
// nothing, which cannot be read.
err := os.Symlink("user@host.1234:1700000000", filepath.Join(dir, ".#"+firstFile))
if err != nil {
t.Fatalf("symlink: %v", err)
}
params, _ := newParams(dir)
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
wantMatched(t, files, get(t, "/probe"), "probe")
}
func TestFaultStopsTheStartNamingTheFileAndLine(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
content string
line int
want string
}{
{
"too few fields", "env-file path ban\n", 1,
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs",
},
{
// Else its regex would be a space, found in nearly every user agent.
"a regex of only spaces and tabs", "scanner user_agent ban\t \n", 1,
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs",
},
{
"an id of other characters", "# ids\n\nenv.file path ban ^/\n", 3,
`the id "env.file" is not an id of letters, digits, - and _`,
},
{
"an unknown target", "env-file paths ban ^/\n", 1,
`the target "paths" is not path, query, uri, method, host, ` +
"user_agent, referer or header:<Name>",
},
{
"a header without a name", "env-file header: ban ^/\n", 1,
`the target "header:" is not path, query, uri, method, host, ` +
"user_agent, referer or header:<Name>",
},
{
"a header name written with its colon", "sqlmap header:User-Agent: ban sqlmap\n", 1,
`the target "header:User-Agent:" has a character after header: ` +
"that no header name can have",
},
{
"a header name with a semicolon", "accept header:Accept;q log ^$\n", 1,
`the target "header:Accept;q" has a character after header: ` +
"that no header name can have",
},
{
"a header name with brackets", "x-header header:X(y) log ^$\n", 1,
`the target "header:X(y)" has a character after header: ` +
"that no header name can have",
},
{
"the Host header", "host-header header:host block ^$\n", 1,
`the target "header:host" names a header that Go's HTTP server ` +
"takes out of every request, so a rule never sees it; " +
"the request's host is the target host",
},
{
"the Transfer-Encoding header",
"# bodies sent in chunks\nchunked header:Transfer-Encoding block ^chunked$\n", 2,
`the target "header:Transfer-Encoding" names a header that Go's ` +
"HTTP server takes out of every request, so a rule never sees it",
},
{
"an unknown action", "env-file path deny ^/\n", 1,
`the action "deny" is not log, block or ban`,
},
{
"a regex that does not compile", "env-file path ban ^/(\n", 1,
"the regex does not compile: error parsing regexp: " +
"missing closing ): `^/(`",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{"00-default.rules": tc.content})
path := filepath.Join(dir, "00-default.rules")
wantRefused(t, dir, path+", line "+strconv.Itoa(tc.line)+": "+tc.want)
})
}
}
func TestIDUsedTwiceStopsTheStartNamingBothPlaces(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{
"00-a.rules": "probe path log ^/a\n",
"50-b.rules": "other path log ^/b\nprobe path ban ^/c\n",
})
wantRefused(t, dir, filepath.Join(dir, "50-b.rules")+`, line 2: the id "probe" `+
"is already the id of the rule at "+filepath.Join(dir, "00-a.rules")+", line 1")
}
func TestDirectoryThatDoesNotExistStopsTheStart(t *testing.T) {
t.Parallel()
dir := filepath.Join(t.TempDir(), "rules.d")
wantRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
": no such file or directory")
}
func TestEmptyDirectoryLoadsNoRulesAndSaysSo(t *testing.T) {
t.Parallel()
params, lines := newParams(writeFiles(t, ruleFiles{"00-default.rules": "# none\n"}))
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
line := lines.waitFor(t, read)
if files.Len() != 0 || line["rules"] != 0.0 {
t.Errorf("%d rules loaded, and the log says %v, want none", files.Len(), line)
}
}
func TestRuleFilesOffReadNothing(t *testing.T) {
t.Parallel()
// SWWAF_RULES_DIR does not exist, which would stop the start.
params, _ := newParams(filepath.Join(t.TempDir(), "rules.d"))
params.Enabled = false
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
if files.Len() != 0 || files.Match(get(t, "/")) != nil {
t.Errorf("%d rules loaded with the rule files off", files.Len())
}
// It would watch until the test ends.
files.Watch(t.Context())
}
func TestEditsTakenInWhileRunning(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines, _ := watch(t, dir)
// matches reports whether path matches a rule.
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
// A file added.
save(t, dir, "50-b.rules", "second path block ^/second\n")
lines.waitUntil(t, func() bool { return matches("/second") })
wantMatched(t, files, get(t, "/first"), "first")
// A file edited.
save(t, dir, firstFile, "first path block ^/edited\n")
lines.waitUntil(t, func() bool { return !matches("/first") })
wantMatched(t, files, get(t, "/edited"), "first")
// A file removed.
err := os.Remove(filepath.Join(dir, "50-b.rules"))
if err != nil {
t.Fatalf("remove: %v", err)
}
lines.waitUntil(t, func() bool { return !matches("/second") })
wantMatched(t, files, get(t, "/edited"), "first")
}
func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines, queue := watch(t, dir)
// The edit's second line has an unknown action, so the rules stay as
// they were, the first line's earlier version included.
save(t, dir, firstFile, "first path block ^/edited\nsecond path bann ^/second\n")
line := lines.waitFor(t, hasError)
want := filepath.Join(dir, firstFile) +
`, line 2: the action "bann" is not log, block or ban`
if line["error"] != want || line["level"] != "ERROR" {
t.Errorf("logged %v, want an error %q", line, want)
}
// The error is raised as a file_error alert too, for the file.
wantFileError := func() {
t.Helper()
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
waiting[0].Reason != hasError || waiting[0].Detail["error"] != want ||
waiting[0].Detail["file"] != filepath.Join(dir, firstFile) {
t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want)
}
}
wantFileError()
wantMatched(t, files, get(t, "/first"), "first")
wantMatched(t, files, get(t, "/second"))
// Once mended, the file is read again, and raises no alert.
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
wantMatched(t, files, get(t, "/edited"), "first")
wantFileError()
}
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
t.Parallel()
params, _ := newParams(filepath.Join("..", "..", "share", "rules.d"))
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load the default file: %v", err)
}
// Probes sent by a browser, by the rule that refuses them.
for rule, targets := range map[string][]string{
"env-file": {"/.env", "/.env.production", "/.ENV"},
"vcs-dir": {"/.git/config", "/.git", "/.svn/entries"},
"secrets-dir": {"/.aws/credentials", "/.ssh/id_rsa"},
"secret-file": {"/.htpasswd", "/.DS_Store", "/.git-credentials"},
"editor-dir": {"/.vscode/sftp.json"},
"backup-file": {
"/wp-config.php.bak", "/index.php~", "/dump.sql", "/backup.sql.gz",
},
"log-file": {"/debug.log"},
"compose-file": {"/docker-compose.yml", "/compose.yaml"},
"php-shell": {"/shell.php"},
"path-traversal": {
"/static/../../etc/passwd", "/f?f=%2e%2e%2f%2e%2e%2fetc%2fpasswd",
},
} {
for _, target := range targets {
wantRefusedBy(t, files, target, browser, rule)
}
}
// Scanners, by their user agents.
for _, scanner := range []string{
"sqlmap/1.8.4#stable (https://sqlmap.org)",
"Mozilla/5.0 (compatible; Nuclei - Open-source project)",
} {
wantRefusedBy(t, files, "/", scanner, "scanner-agent")
}
// Ordinary requests to a code forge for files of those names deeper
// in its paths, and for other files at its root.
for _, target := range []string{
"/owner/repo/src/branch/main/.env.example",
"/owner/repo/src/branch/main/.env",
"/owner/repo/src/branch/main/.github/workflows/ci.yml",
"/owner/repo/src/branch/main/.vscode/settings.json",
"/owner/repo/src/branch/main/.htaccess",
"/owner/repo/src/branch/main/docker-compose.yml",
"/owner/repo/src/branch/main/db/schema.sql",
"/owner/repo/raw/branch/main/debug.log",
"/owner/repo.git/info/refs?service=git-upload-pack",
"/owner/repo/src/branch/main/docs/../README.md",
"/user/login?redirect_to=%2fowner%2frepo",
"/index.php",
"/.well-known/security.txt",
} {
r := get(t, target)
r.Header.Set(userAgent, browser)
matched := files.Match(r)
if len(matched) != 0 {
t.Errorf("%s matched %v, want no rule", target, ids(matched))
}
}
// A request without a user agent is only noted.
wantMatched(t, files, get(t, "/"), "empty-agent")
}
// ruleFiles are files to write into a directory of rule files, by name.
type ruleFiles map[string]string
// writeFiles writes files into a new directory, and returns it.
func writeFiles(t *testing.T, files ruleFiles) string {
t.Helper()
dir := t.TempDir()
for name, content := range files {
path := filepath.Join(dir, name)
err := os.MkdirAll(filepath.Dir(path), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
err = os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
}
return dir
}
// save writes content to the rule file name in dir as an editor that
// saves by renaming does, so that the file is never seen half written.
func save(t *testing.T, dir, name, content string) {
t.Helper()
path := filepath.Join(dir, name)
err := os.WriteFile(path+".tmp", []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
err = os.Rename(path+".tmp", path)
if err != nil {
t.Fatalf("rename: %v", err)
}
}
// newParams returns Params for the rule files in dir, switched on, with
// the process log in the processLog returned, and the alerts waiting in a
// queue for a webhook that is never sent them.
func newParams(dir string) (rules.Params, processLog) {
lines := make(processLog, maxLogLines)
return rules.Params{
Dir: dir,
Enabled: true,
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
Alerts: alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: time.Now,
}),
}, lines
}
// load writes files into a new directory and loads the rules in it.
func load(t *testing.T, files ruleFiles) *rules.Files {
t.Helper()
params, _ := newParams(writeFiles(t, files))
params.ProcessLog = slog.New(slog.DiscardHandler)
loaded, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
return loaded
}
// watch loads the rules in dir, runs their Watch until the test ends, and
// waits until it watches the directory. It returns the alerts' queue as
// well.
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
t.Helper()
params, lines := newParams(dir)
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Watch(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
lines.waitFor(t, watching)
return files, lines, params.Alerts
}
// wantRefused checks that loading the rule files in dir fails with the
// error want.
func wantRefused(t *testing.T, dir, want string) {
t.Helper()
params, _ := newParams(dir)
_, err := rules.Load(params)
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
// get returns a GET request for target, a path and an optional query, as
// smallwebwaf's server reads it, without a user agent.
func get(t *testing.T, target string) *http.Request {
t.Helper()
return httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"http://app.example"+target, nil)
}
// wantRefusedBy checks that a GET request for target with the user agent
// sent matches rule alone, and that rule refuses it.
func wantRefusedBy(t *testing.T, files *rules.Files, target, sent, rule string) {
t.Helper()
r := get(t, target)
r.Header.Set(userAgent, sent)
matched := files.Match(r)
if len(matched) != 1 || matched[0].ID != rule || matched[0].Action == rules.ActionLog {
t.Errorf("%s from %q matched %v, want %s alone, refusing it", target,
sent, ids(matched), rule)
}
}
// wantMatched checks the ids of the rules r matches, in order.
func wantMatched(t *testing.T, files *rules.Files, r *http.Request, want ...string) {
t.Helper()
got := ids(files.Match(r))
if !slices.Equal(got, want) {
t.Errorf("%s matched %v, want %v", r.URL, got, want)
}
}
// ids returns the ids of matched.
func ids(matched []rules.Rule) []string {
got := make([]string, 0, len(matched))
for _, rule := range matched {
got = append(got, rule.ID)
}
return got
}
// processLog receives the lines of a process log, each a JSON object, for
// a test to wait for.
type processLog chan string
// Write receives a line of the process log.
func (l processLog) Write(line []byte) (int, error) {
l <- string(line)
return len(line), nil
}
// waitFor returns the next line of the process log whose message is msg,
// passing over the lines before it. It waits as long as that takes, so
// that a slow test process cannot fail the test.
func (l processLog) waitFor(t *testing.T, msg string) map[string]any {
t.Helper()
for line := range l {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil {
t.Fatalf("process log line %q is not JSON: %v", line, err)
}
if fields["msg"] == msg {
return fields
}
}
return nil
}
// waitUntil waits for the rule files to be read until done reports true,
// as it does once they have been read after the test's last change. They
// can be read before then too, as they are once Watch starts watching.
func (l processLog) waitUntil(t *testing.T, done func() bool) {
t.Helper()
for !done() {
l.waitFor(t, read)
}
}
+165
View File
@@ -0,0 +1,165 @@
package rules
import (
"context"
"log/slog"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"slices"
"testing"
"testing/synctest"
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// The tests below run readAfterChanges in a synctest bubble, where time is
// a clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once readAfterChanges waits again, so that every
// reading due by then is done. The test sends the changes itself, as the
// watch of a directory cannot run in a bubble.
func TestFileWrittenInTwoPartsTakenInOnlyWhole(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "50-app.rules")
writeFile(t, path, "first path block ^/first\n")
files := load(t, dir)
changes := run(t, files)
file, err := os.Create(path) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("create: %v", err)
}
defer func() {
_ = file.Close()
}()
// The first part ends in the middle of a ban rule's regex, which,
// read then, would ban every request.
write(t, file, "first path block ^/first\nprobe path ban ^/")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/anything")
// The second part starts the wait again.
write(t, file, `\.env$`+"\n")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/.env")
time.Sleep(time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/.env", "probe")
wantMatched(t, files, "/anything")
})
}
func TestEditSavedBeforeTheWatchStartsTakenIn(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "50-app.rules")
writeFile(t, path, "first path block ^/first\n")
files := load(t, dir)
// Saved after Load read the files, and before the directory was
// watched, so that no change is seen for it.
writeFile(t, path, "first path block ^/edited\n")
run(t, files)
time.Sleep(quietTime)
synctest.Wait()
wantMatched(t, files, "/edited", "first")
})
}
// load loads the rules in dir.
func load(t *testing.T, dir string) *Files {
t.Helper()
files, err := Load(Params{
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
Alerts: alerts.New(alerts.Params{}),
})
if err != nil {
t.Fatalf("load: %v", err)
}
return files
}
// run runs files' readAfterChanges until the test ends, and returns the
// channel that sends it changes.
func run(t *testing.T, files *Files) chan<- fsnotify.Event {
t.Helper()
changes := make(chan fsnotify.Event)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.readAfterChanges(ctx, changes, nil)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
return changes
}
// writeFile writes content to the file at path.
func writeFile(t *testing.T, path, content string) {
t.Helper()
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
}
// write writes text to the end of file.
func write(t *testing.T, file *os.File, text string) {
t.Helper()
_, err := file.WriteString(text)
if err != nil {
t.Fatalf("write: %v", err)
}
}
// wantMatched checks the ids of the rules that a GET request for path
// matches, in order.
func wantMatched(t *testing.T, files *Files, path string, want ...string) {
t.Helper()
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"http://app.example"+path, nil)
matched := files.Match(r)
got := make([]string, 0, len(matched))
for _, rule := range matched {
got = append(got, rule.ID)
}
if !slices.Equal(got, want) {
t.Errorf("%s matched %v, want %v", path, got, want)
}
}
+5 -3
View File
@@ -23,6 +23,8 @@ var errHealthEndpoint = errors.New("smallwebwaf's health endpoint answered")
// smallwebwaf answers its health endpoint on 127.0.0.1, at the port in // smallwebwaf answers its health endpoint on 127.0.0.1, at the port in
// SWWAF_LISTEN_ADDR, and the app accepts connections at the address in // SWWAF_LISTEN_ADDR, and the app accepts connections at the address in
// SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1. // SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1.
// It reads no other setting, nor a file that another names, so neither
// can fail it.
// args are the arguments after `healthcheck`; it takes none, and given // args are the arguments after `healthcheck`; it takes none, and given
// one it names it on stderr and returns 1 without checking anything. // one it names it on stderr and returns 1 without checking anything.
func HealthCheck( func HealthCheck(
@@ -50,13 +52,13 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout) ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout)
defer cancel() defer cancel()
cfg, err := config.FromEnvironment(lookupEnv) listenAddr, upstreamURL, err := config.ListenAddrAndUpstreamURL(lookupEnv)
if err != nil { if err != nil {
return fmt.Errorf("invalid setting: %w", err) return fmt.Errorf("invalid setting: %w", err)
} }
// The settings have checked that the address has a port. // The settings have checked that the address has a port.
_, port, _ := net.SplitHostPort(cfg.ListenAddr) _, port, _ := net.SplitHostPort(listenAddr)
health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath
req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody) req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody)
@@ -75,7 +77,7 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
return fmt.Errorf("%w %s", errHealthEndpoint, res.Status) return fmt.Errorf("%w %s", errHealthEndpoint, res.Status)
} }
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(cfg.UpstreamURL)) conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(upstreamURL))
if err != nil { if err != nil {
return fmt.Errorf("connect to the app: %w", err) return fmt.Errorf("connect to the app: %w", err)
} }
+39 -4
View File
@@ -6,6 +6,8 @@ import (
"net" "net"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"os"
"path/filepath"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -24,12 +26,15 @@ func TestHealthCheck(t *testing.T) {
out := &output{} out := &output{}
exited := make(chan int, 1) exited := make(chan int, 1)
settings := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
}
go func() { go func() {
exited <- run(ctx, map[string]string{ exited <- run(ctx, settings, out)
listenAddr: localhost + ":0",
upstreamURL: app.URL,
}, out)
}() }()
addr, _ := out.line(t, "msg", "starting")["address"].(string) addr, _ := out.line(t, "msg", "starting")["address"].(string)
@@ -40,6 +45,21 @@ func TestHealthCheck(t *testing.T) {
wantHealthCheck(t, env, 0, "") wantHealthCheck(t, env, 0, "")
// The health check reads those two settings alone, here given as
// files: a removed or invalid token file, or an invalid value of
// another setting, does not fail it.
for _, other := range []struct{ name, value string }{
{"SWWAF_METRICS_TOKEN_FILE", filepath.Join(t.TempDir(), "removed")},
{"SWWAF_METRICS_TOKEN_FILE", writeFile(t, "too short\n")},
{"SWWAF_MODE", "neither"},
} {
wantHealthCheck(t, map[string]string{
listenAddr + "_FILE": writeFile(t, ":"+port+"\n"),
upstreamURL + "_FILE": writeFile(t, app.URL+"\n"),
other.name: other.value,
}, 0, "")
}
app.Close() app.Close()
wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ") wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ")
@@ -95,3 +115,18 @@ func wantHealthCheck(t *testing.T, env map[string]string, status int, message st
got, wrote, status, message) got, wrote, status, message)
} }
} }
// writeFile writes contents to a file in a directory of its own, removed
// when the test ends, and returns the file's path.
func writeFile(t *testing.T, contents string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "setting")
err := os.WriteFile(path, []byte(contents), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
return path
}
+193 -15
View File
@@ -1,6 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings, // Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// serves requests until it is told to stop, and then stops in an orderly // the rule files and the state files, serves requests until it is told to
// way. // stop, and then stops in an orderly way, writing the state files.
package smallwebwaf package smallwebwaf
import ( import (
@@ -15,10 +15,14 @@ import (
"syscall" "syscall"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
"sneak.berlin/go/smallwebwaf/internal/state"
) )
// shutdownTimeout is how long requests in progress may take to finish // shutdownTimeout is how long requests in progress may take to finish
@@ -26,6 +30,11 @@ import (
// runit and docker wait a little longer before they kill the process. // runit and docker wait a little longer before they kill the process.
const shutdownTimeout = 5 * time.Second const shutdownTimeout = 5 * time.Second
// remoteLogStopTimeout is how long, as smallwebwaf stops, the log lines
// still waiting are sent to SWWAF_LOG_REMOTE_URL before they are given
// up. stdout has carried them.
const remoteLogStopTimeout = 2 * time.Second
// Params are what Run needs from the process. // Params are what Run needs from the process.
type Params struct { type Params struct {
// Version is the version of the binary, set when it is built. // Version is the version of the binary, set when it is built.
@@ -55,8 +64,9 @@ func Main(version string) int {
}) })
} }
// Run reads the settings, then serves requests until ctx is done. It // Run reads the settings, the rule files and the state files, then serves
// returns the process's exit status, 1 when smallwebwaf cannot start. // requests until ctx is done. It returns the process's exit status, 1
// when smallwebwaf cannot start.
func Run(ctx context.Context, params Params) int { func Run(ctx context.Context, params Params) int {
processLog := requestlog.NewProcessLogger(params.Stdout) processLog := requestlog.NewProcessLogger(params.Stdout)
@@ -67,6 +77,60 @@ func Run(ctx context.Context, params Params) int {
return 1 return 1
} }
// While SWWAF_LOG_REMOTE_URL is set, every line on stdout from here on
// is sent there too.
stdout := params.Stdout
var remote *remotelog.Sender
if cfg.LogRemoteURL != nil {
remote = newRemoteLogSender(cfg)
stdout = io.MultiWriter(params.Stdout, remote)
processLog = requestlog.NewProcessLogger(stdout)
stopSending := startSending(ctx, remote, processLog)
defer stopSending()
}
// The state files and the alerts give times in UTC.
now := func() time.Time { return time.Now().UTC() }
alertQueue := newAlertQueue(cfg, now, processLog)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir,
Enabled: cfg.RulesEnabled,
ProcessLog: processLog,
Alerts: alertQueue,
})
if err != nil {
processLog.Error("cannot use the rule files", "error", err.Error())
return 1
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
})
if remote != nil {
server.Metrics.AddRemoteLog(remote)
}
server.Metrics.AddAlerts(alertQueue)
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
if err != nil {
processLog.Error("cannot use the state files", "error", err.Error())
return 1
}
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr) listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
if err != nil { if err != nil {
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR", processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
@@ -75,26 +139,99 @@ func Run(ctx context.Context, params Params) int {
return 1 return 1
} }
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: time.Now,
})
processLog.Info("starting", processLog.Info("starting",
"version", params.Version, "version", params.Version,
"address", listener.Addr().String(), "address", listener.Addr().String(),
"settings", cfg) "settings", cfg)
return serve(ctx, server, listener, processLog) return serve(ctx, server.Server, listener, files, ruleFiles, alertQueue, processLog)
} }
// serve serves requests on listener until ctx is done, then gives the // loadStateFiles reads the state files into the parts of server and into
// requests in progress shutdownTimeout to finish. // alertQueue, as state.Load does.
func loadStateFiles(
cfg *config.Config, server *proxy.Server, alertQueue *alerts.Queue,
now func() time.Time, processLog *slog.Logger,
) (*state.Files, error) {
return state.Load(state.Params{
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Alerts: alertQueue,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
}
// newAlertQueue returns the queue of the alerts to the webhook, Slack and
// ntfy, with the settings for them.
func newAlertQueue(
cfg *config.Config, now func() time.Time, processLog *slog.Logger,
) *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
WebhookHeaders: cfg.AlertWebhookHeaders,
SlackURL: cfg.AlertSlackWebhookURL,
NtfyURL: cfg.AlertNtfyURL,
NtfyToken: cfg.AlertNtfyToken,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
}
// newRemoteLogSender returns a sender of the log lines to
// SWWAF_LOG_REMOTE_URL, with the settings for it.
func newRemoteLogSender(cfg *config.Config) *remotelog.Sender {
return remotelog.New(remotelog.Params{
URL: cfg.LogRemoteURL,
RootCAs: cfg.LogRemoteTLSCAs,
Buffer: cfg.LogRemoteBuffer,
Facility: cfg.LogRemoteFacility,
AppName: cfg.LogRemoteAppName,
})
}
// startSending runs remote until the function it returns is called, which
// then waits at most remoteLogStopTimeout for the lines still waiting to
// be sent. Sending goes on after ctx is done, so that the lines written
// while smallwebwaf stops are sent too.
func startSending(
ctx context.Context, remote *remotelog.Sender, processLog *slog.Logger,
) func() {
sending, stop := context.WithCancel(context.WithoutCancel(ctx))
sent := make(chan struct{})
go func() {
remote.Run(sending, processLog)
close(sent)
}()
return func() {
stop()
select {
case <-sent:
case <-time.After(remoteLogStopTimeout):
}
}
}
// 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
// they change, 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 the alerts still waiting.
func serve( func serve(
ctx context.Context, server *http.Server, listener net.Listener, ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
processLog *slog.Logger, processLog *slog.Logger,
) int { ) int {
served := make(chan error, 1) served := make(chan error, 1)
@@ -103,6 +240,14 @@ func serve(
served <- server.Serve(listener) served <- server.Serve(listener)
}() }()
writing, stopWriting := context.WithCancel(ctx)
defer stopWriting()
written := inBackground(func() { files.Run(writing) })
watched := inBackground(func() { files.Watch(writing) })
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
alertsSent := inBackground(func() { alertQueue.Run(writing) })
select { select {
case err := <-served: case err := <-served:
processLog.Error("serving failed", "error", err.Error()) processLog.Error("serving failed", "error", err.Error())
@@ -132,7 +277,40 @@ func serve(
return 1 return 1
} }
// Run and Watch have ended, so nothing else reads or writes the
// files, and no alert is being sent, so that alerts.json keeps every
// alert not yet sent. Every request has ended too, but for two kinds
// that Go's server does not wait for: one cut off because Shutdown
// timed out, and one whose connection switched protocols, such as a
// WebSocket. Such a request adds to its client's history only as it
// ends, which can be after this write, and then that request is
// missing from clients.json.
<-written
<-watched
<-rulesWatched
<-alertsSent
err = files.WriteAll()
if err != nil {
processLog.Error("writing the state files failed", "error", err.Error())
return 1
}
processLog.Info("stopped") processLog.Info("stopped")
return 0 return 0
} }
// inBackground runs task on a goroutine of its own, and returns a channel
// that is closed once task has returned.
func inBackground(task func()) <-chan struct{} {
done := make(chan struct{})
go func() {
task()
close(done)
}()
return done
}
@@ -0,0 +1,82 @@
package smallwebwaf
import (
"log/slog"
"net/url"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// The stop's tests run in a synctest bubble, where the time package runs
// on a clock of the test's own, so that how long the stop takes can be
// told exactly. The sender is held up by its process log, not by the
// network: a goroutine of the bubble that waits on the network keeps that
// clock from moving on.
func TestStopWaitsForTheSenderToFinish(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
took := stopHeldSender(t, time.Second)
if took != time.Second {
t.Errorf("the stop took %s, want the second the sender took", took)
}
})
}
func TestStopWaitsForTheSenderAtMostTwoSeconds(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
took := stopHeldSender(t, time.Minute)
if took != 2*time.Second {
t.Errorf("the stop took %s, want 2s", took)
}
})
}
// heldLog holds each line written to it until it is closed.
type heldLog chan struct{}
// Write waits until the log is closed.
func (l heldLog) Write(p []byte) (int, error) {
<-l
return len(p), nil
}
// stopHeldSender starts sending to an endpoint the sender cannot connect
// to, holds the sender as it logs that failure until release has passed,
// stops the sending, and returns how long the stop took. It returns once
// the sender has ended, as a bubble must.
func stopHeldSender(t *testing.T, release time.Duration) time.Duration {
t.Helper()
log := make(heldLog)
sender := remotelog.New(remotelog.Params{
// No port is 65536, so each attempt to connect fails at once,
// before it reaches the network.
URL: &url.URL{Scheme: remotelog.SchemeTCP, Host: "127.0.0.1:65536"},
Buffer: 1,
})
stopSending := startSending(t.Context(), sender,
slog.New(slog.NewJSONHandler(log, nil)))
synctest.Wait()
time.AfterFunc(release, func() { close(log) })
stopped := time.Now()
stopSending()
took := time.Since(stopped)
time.Sleep(release)
synctest.Wait()
return took
}
File diff suppressed because it is too large Load Diff
+851
View File
@@ -0,0 +1,851 @@
// Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and
// history, lookups.json GeoJS's answers, and alerts.json the cooldowns,
// the hour under way and the alerts waiting for each destination. Load
// reads them at start, Watch takes in an admin's edit of one while
// smallwebwaf runs, and Run and WriteAll write them. The disk is read and
// written outside the parts' locks, which are held only to take a
// snapshot or to put in what a file holds, so that no request waits on
// the disk.
package state
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io/fs"
"log/slog"
"maps"
"net/netip"
"os"
"path/filepath"
"slices"
"sync"
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// version is the version of the files' format, the only one read.
const version = 1
// fileMode lets the smallwebwaf user alone read and write the files, which
// hold visitors' addresses.
const fileMode = 0o600
// The state files' names.
const (
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
alertsJSON = "alerts.json"
)
var (
errVersion = errors.New("unknown version")
// errMissing is for an entry without a field it needs.
errMissing = errors.New("has no")
errCause = errors.New("is not limit, attack or admin")
errDestination = errors.New("is not webhook, slack or ntfy")
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
`or remove the file`)
)
// Params are what Load needs.
type Params struct {
// Dir is the directory of the state files (SWWAF_STATE_DIR).
Dir string
// WriteDelay is how long after a ban is made bans.json is written
// (SWWAF_STATE_WRITE_DELAY), and CounterInterval how often every file
// is (SWWAF_STATE_COUNTER_INTERVAL).
WriteDelay time.Duration
CounterInterval time.Duration
// Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
// receive a file_error alert for an edit set aside, and for a write
// that fails while smallwebwaf runs.
Ledger *bans.Ledger
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
Alerts *alerts.Queue
// Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC.
Now func() time.Time
// ProcessLog receives what was read and taken in, the edits set aside,
// and the writes that fail.
ProcessLog *slog.Logger
// Metrics count each file's writes, and the edits taken in and set
// aside.
Metrics *metrics.Metrics
}
// Files are the state files of a running smallwebwaf.
type Files struct {
params Params
// mu is held while a file is read for an edit, and while it is
// written, so that Watch and the writes take turns. No request takes
// it.
mu sync.Mutex
// sums are the SHA-256 sums of what each file held, by name, when
// smallwebwaf last read or wrote it. A file that holds anything else
// has been edited since.
sums map[string][sha256.Size]byte
}
// bansFile is bans.json, indented for an admin to read and edit.
type bansFile struct {
Version int `json:"version"`
Bans []BanEntry `json:"bans"`
}
// BanEntry is a ban as bans.json holds it: a permanent ban's expires is
// null, a ban an admin added may have no cause, which makes it an
// admin's, and lifted is left out until an admin lifts the ban. The ban
// endpoints answer with bans in this form too.
type BanEntry struct {
Netblock netip.Prefix `json:"netblock"`
Start time.Time `json:"start"`
Expires *time.Time `json:"expires"`
Cause string `json:"cause"`
Reason string `json:"reason,omitempty"`
Lifted *time.Time `json:"lifted,omitempty"`
Notes bans.Notes `json:"notes"`
}
// clientsFile is clients.json, with each client on a line of its own.
type clientsFile struct {
Version int `json:"version"`
Clients []ratelimit.Client `json:"clients"`
}
// lookupsFile is lookups.json, with each answer on a line of its own.
type lookupsFile struct {
Version int `json:"version"`
Lookups []lookup.Answer `json:"lookups"`
}
// alertsFile is alerts.json, indented for an admin to read and edit.
type alertsFile struct {
Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"`
Waiting map[string][]alerts.Alert `json:"waiting"`
}
// stateFile is the struct of a state file. Once the file is decoded, its
// check refuses the first entry without a field it needs, which would
// otherwise be read as something the entry does not say. data is the
// file, for a field that may be null or "" but not left out, which the
// struct cannot tell apart.
type stateFile interface {
check(data []byte) error
}
// 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
// state, as on a first start. A file that does not parse, has an unknown
// version, or has an entry without a field it needs, is an error that
// names the file and, where the JSON decoder tells it, the line and
// column, or else the entry.
func Load(params Params) (*Files, error) {
err := checkWritable(params.Dir)
if err != nil {
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
}
f := &Files{params: params, sums: map[string][sha256.Size]byte{}}
bansRead, bansErr := f.read(bansJSON)
clientsRead, clientsErr := f.read(clientsJSON)
lookupsRead, lookupsErr := f.read(lookupsJSON)
alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
if err != nil {
return nil, err
}
params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
"alerts_waiting", alertsRead)
return f, nil
}
// Run writes bans.json WriteDelay after a ban is made, with every ban
// made in between, and every file every CounterInterval, until ctx is
// done. A write that fails is logged, raised as a file_error alert, and
// the file is written again at its next write. Each write takes in an
// admin's edit of its file first, as writeFile describes.
func (f *Files) Run(ctx context.Context) {
interval := time.NewTicker(f.params.CounterInterval)
defer interval.Stop()
var bansDue <-chan time.Time // nil while no ban waits to be written
for {
select {
case <-ctx.Done():
return
case <-f.params.Ledger.Changed():
if bansDue == nil {
bansDue = time.After(f.params.WriteDelay)
}
case <-bansDue:
bansDue = nil
f.logFailure(bansJSON, f.writeFile(bansJSON))
case <-interval.C:
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
f.logFailure(name, f.writeFile(name))
}
}
}
}
// WriteAll writes every state file, as smallwebwaf stops. A file that
// fails does not keep the others from being written.
func (f *Files) WriteAll() error {
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON), f.writeFile(alertsJSON))
}
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
// state file as soon as it is saved: what the file holds replaces what
// smallwebwaf held for it. An edit that does not parse is left for the
// file's next write, which sets it aside, since a file can be read while
// an editor is still writing it. If Dir cannot be watched, that is
// logged, and an edit is taken in only before its file is written.
func (f *Files) Watch(ctx context.Context) {
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the state files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the state files for edits",
"directory", f.params.Dir)
for {
select {
case <-ctx.Done():
return
case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
f.fileChanged(name)
}
case err = <-watcher.Errors:
f.params.ProcessLog.Warn("watching the state files failed",
"error", err.Error())
}
}
}
// logFailure logs a write of the state file name that failed, and raises
// a file_error alert for it.
func (f *Files) logFailure(name string, err error) {
if err != nil {
const failed = "writing the state files failed"
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: failed,
Detail: map[string]any{
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
},
})
f.params.ProcessLog.Error(failed, "error", err.Error())
}
}
// fileChanged takes in what the state file name holds, as Watch sees it
// change, if that is an edit made since smallwebwaf last read or wrote
// the file. A file that cannot be read or does not parse is left for its
// next write.
func (f *Files) fileChanged(name string) {
f.mu.Lock()
defer f.mu.Unlock()
data, changed, err := f.readChanged(name)
if err != nil || !changed {
return
}
_ = f.takeInEdit(name, data)
}
// takeInEdit takes in data, an edit of the state file name, as takeIn
// does, and counts and logs it. Every edit taken in while smallwebwaf
// runs, by Watch or by a write, is taken in here. An edit that does not
// parse is neither counted nor logged, and takeIn's error returned.
func (f *Files) takeInEdit(name string, data []byte) error {
_, err := f.takeIn(name, data, true)
if err != nil {
return err
}
// Counted before it is logged, so that the count is there once the
// log line is.
f.params.Metrics.StateFileEditTakenIn(name)
f.params.ProcessLog.Info("took in an edit of a state file",
"file", filepath.Join(f.params.Dir, name))
return nil
}
// read takes in the state file name at start, and returns how many
// entries it holds. A missing file holds none.
func (f *Files) read(name string) (int, error) {
data, changed, err := f.readChanged(name)
if err != nil || !changed {
return 0, err
}
return f.takeIn(name, data, false)
}
// readChanged returns what the state file name holds, and whether that
// has changed since smallwebwaf last read or wrote the file, as it has
// for a file smallwebwaf never read or wrote. A missing file has not
// changed: it is written again at its next write.
func (f *Files) readChanged(name string) ([]byte, bool, error) {
path := filepath.Join(f.params.Dir, name)
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
if errors.Is(err, fs.ErrNotExist) {
return nil, false, nil
}
if err != nil {
return nil, false, err
}
return data, sha256.Sum256(data) != f.sums[name], nil
}
// takeIn parses data, what the state file name holds, puts it into the
// part that keeps that state, in place of what the part held, and returns
// how many entries the file holds. edit is whether data is an admin's
// edit taken in while smallwebwaf runs, rather than the file read at the
// start. An error names the file and, where the JSON decoder tells it,
// the line and column, or else the entry.
func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
path := filepath.Join(f.params.Dir, name)
var entries int
switch name {
case bansJSON:
var file bansFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
held := make([]bans.Ban, 0, len(file.Bans))
for _, entry := range file.Bans {
held = append(held, entry.ban())
}
if edit {
f.params.Ledger.LoadEdit(held)
} else {
f.params.Ledger.Load(held)
}
entries = len(held)
case clientsJSON:
var file clientsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Limiter.Load(file.Clients, f.params.Now())
entries = len(file.Clients)
case lookupsJSON:
var file lookupsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups)
case alertsJSON:
// waiting was a list, of the alerts waiting for the webhook, before
// alerts went to Slack and ntfy too.
var written struct {
Waiting json.RawMessage `json:"waiting"`
}
if json.Unmarshal(data, &written) == nil &&
bytes.HasPrefix(written.Waiting, []byte("[")) {
return 0, fmt.Errorf("%s: %w", path, errWaitingList)
}
var file alertsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Alerts.Load(alerts.State{
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
})
for _, waiting := range file.Waiting {
entries += len(waiting)
}
}
f.sums[name] = sha256.Sum256(data)
return entries, nil
}
// writeFile writes the state file name from what smallwebwaf holds. An
// edit made since smallwebwaf last read or wrote the file is taken in
// first, so that it is not overwritten, or set aside if it does not
// parse. A file that cannot be read, or an edit that cannot be set
// aside, is left as it is, and the write given up. Every write is counted
// in the metrics, and one that fails or is given up as a failure.
func (f *Files) writeFile(name string) error {
f.mu.Lock()
defer f.mu.Unlock()
data, changed, err := f.readChanged(name)
if err == nil && changed {
err = f.takeInEdit(name, data)
if err != nil {
err = f.setAside(name, err)
}
}
if err == nil {
data, err = f.encode(name)
if err != nil {
err = fmt.Errorf("encode %s: %w", name, err)
}
}
if err == nil {
err = write(f.params.Dir, name, data)
}
if err == nil {
// The file holds data from here on, even if the directory sync
// fails, so that its next read does not take it for an admin's
// edit.
f.sums[name] = sha256.Sum256(data)
err = syncDirectory(f.params.Dir)
}
f.params.Metrics.StateFileWritten(name, len(data), err)
return err
}
// setAside renames the state file name, an edit that does not parse with
// parseErr, to name.bad, for the admin to mend, and logs it with where in
// the file the error is, and raises a file_error alert for it. If the
// rename fails, the edit is left as it is, and the error returned is
// parseErr joined with the rename's.
func (f *Files) setAside(name string, parseErr error) error {
path := filepath.Join(f.params.Dir, name)
err := os.Rename(path, path+".bad")
if err != nil {
return errors.Join(parseErr, err)
}
const setAside = "set aside an edit of a state file that does not parse"
// Raised before it is logged, so that the alert is there once the log
// line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: setAside,
Detail: map[string]any{"file": path + ".bad", "error": parseErr.Error()},
})
f.params.ProcessLog.Error(setAside, "file", path+".bad", "error", parseErr.Error())
f.params.Metrics.StateFileEditSetAside(name)
return nil
}
// encode returns the state file name as smallwebwaf writes it, from a
// snapshot of the part that keeps that state.
func (f *Files) encode(name string) ([]byte, error) {
switch name {
case bansJSON:
file := 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:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
case lookupsJSON:
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
default: // alerts.json
held := f.params.Alerts.Snapshot()
file := alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting,
}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
}
}
// BanEntries returns held as bans.json lists them, an empty list for
// none.
func BanEntries(held []bans.Ban) []BanEntry {
entries := make([]BanEntry, 0, len(held))
for _, ban := range held {
entries = append(entries, newBanEntry(ban))
}
return entries
}
// newBanEntry returns ban as bans.json holds it.
func newBanEntry(ban bans.Ban) BanEntry {
entry := BanEntry{
Netblock: ban.Netblock, Start: ban.Start, Cause: ban.Cause, Reason: ban.Reason,
Notes: ban.Notes,
}
if !ban.Permanent() {
entry.Expires = &ban.Expires
}
if !ban.Lifted.IsZero() {
entry.Lifted = &ban.Lifted
}
return entry
}
// ban returns the ban an entry of bans.json holds.
func (e BanEntry) ban() bans.Ban {
ban := bans.Ban{
Netblock: e.Netblock, Start: e.Start, Cause: e.Cause, Reason: e.Reason,
Notes: e.Notes,
}
if e.Expires != nil {
ban.Expires = *e.Expires
}
if e.Lifted != nil {
ban.Lifted = *e.Lifted
}
return ban
}
// check refuses a ban without a netblock, which would refuse every IPv6
// client, a start, from which the length of the netblock's next ban is
// 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
// each expires is read again as written. A cause other than limit,
// attack or admin, most likely misspelt, is refused too.
func (f *bansFile) check(data []byte) error {
var written struct {
Bans []struct {
Expires json.RawMessage `json:"expires"`
} `json:"bans"`
}
err := json.Unmarshal(data, &written)
if err != nil {
return err
}
for i, entry := range f.Bans {
switch {
case !entry.Netblock.IsValid():
return missing(i, "netblock")
case entry.Start.IsZero():
return missing(i, "start")
case written.Bans[i].Expires == nil:
return missing(i, "expires")
case entry.Cause != "" && entry.Cause != bans.CauseLimit &&
entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin:
return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause)
}
}
return nil
}
// 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
// them and give the client a fresh allowance.
func (f *clientsFile) check([]byte) error {
for i, client := range f.Clients {
switch {
case !client.Client.IsValid():
return missing(i, "client")
case countsWithoutStart(client.Minute):
return missing(i, "minute.start")
case countsWithoutStart(client.Hour):
return missing(i, "hour.start")
case countsWithoutStart(client.Day):
return missing(i, "day.start")
}
}
return nil
}
// check refuses an answer without a client, which would answer for
// nobody, a country, which would place the client nowhere, or the time
// GeoJS gave it, which would drop it. "" is the country of a client
// GeoJS cannot place, which Lookups cannot tell from a missing one, so
// each country is read again as written.
func (f *lookupsFile) check(data []byte) error {
var written struct {
Lookups []struct {
Country *string `json:"country"`
} `json:"lookups"`
}
err := json.Unmarshal(data, &written)
if err != nil {
return err
}
for i, answer := range f.Lookups {
switch {
case !answer.Client.IsValid():
return missing(i, "client")
case written.Lookups[i].Country == nil:
return missing(i, "country")
case answer.Answered.IsZero():
return missing(i, "answered")
}
}
return nil
}
// check refuses a cooldown without its event or when its alert was sent,
// which would hold back no repeat, alerts waiting for a destination with
// another name than webhook, slack or ntfy, most likely misspelt, and an
// alert waiting without its event or its time.
func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns {
switch {
case cooldown.Event == "":
return fmt.Errorf("cooldowns %w", missing(i, "event"))
case cooldown.Sent.IsZero():
return fmt.Errorf("cooldowns %w", missing(i, "sent"))
}
}
for _, destination := range slices.Sorted(maps.Keys(f.Waiting)) {
if !slices.Contains(alerts.Destinations(), destination) {
return fmt.Errorf("waiting %q %w", destination, errDestination)
}
for i, alert := range f.Waiting[destination] {
switch {
case alert.Event == "":
return fmt.Errorf("waiting %s %w", destination, missing(i, "event"))
case alert.Time.IsZero():
return fmt.Errorf("waiting %s %w", destination, missing(i, "time"))
}
}
}
return nil
}
// countsWithoutStart reports whether b holds requests but no start, which
// places them in time.
func countsWithoutStart(b ratelimit.Buckets) bool {
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
}
// missing returns the error for entry i, counted from 0, of a state file,
// which has no field.
func missing(i int, field string) error {
return fmt.Errorf("entry %d %w %q", i+1, errMissing, field)
}
// encodeOnePerLine encodes a state file whose entries, under key, are one
// to a line, so that grep shows everything about one client.
func encodeOnePerLine[E any](key string, entries []E) ([]byte, error) {
var b bytes.Buffer
fmt.Fprintf(&b, "{\n \"version\": %d,\n %q: [", version, key)
for i, entry := range entries {
line, err := json.Marshal(entry)
if err != nil {
return nil, err
}
if i > 0 {
b.WriteString(",")
}
b.WriteString("\n ")
b.Write(line)
}
b.WriteString("\n ]\n}\n")
return b.Bytes(), nil
}
// checkWritable makes a file in dir and removes it again.
func checkWritable(dir string) error {
file, err := os.CreateTemp(dir, "write-check-*")
if err != nil {
return err
}
return errors.Join(file.Close(), os.Remove(file.Name()))
}
// parse reads data, what the state file at path holds, into file, a
// pointer to that file's struct, and checks its entries.
func parse(path string, data []byte, file stateFile) error {
// The version is read first, so that a file of another version is
// refused for that, and not for an entry this version cannot read.
var header struct {
Version int `json:"version"`
}
err := json.Unmarshal(data, &header)
if err == nil && header.Version != version {
err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d",
errVersion, header.Version, version)
}
if err == nil {
decoder := json.NewDecoder(bytes.NewReader(data))
// A field this version does not know is most likely misspelt, and
// its value would be lost without a word.
decoder.DisallowUnknownFields()
err = decoder.Decode(file)
}
if err == nil {
err = file.check(data)
}
if err != nil {
return fmt.Errorf("%s%s: %w", path, position(data, err), err)
}
return nil
}
// position returns where in data err was found, as ", line L, column C"
// of the last byte the JSON decoder read, or "" when err does not tell.
func position(data []byte, err error) string {
var (
syntaxErr *json.SyntaxError
typeErr *json.UnmarshalTypeError
read int64
)
switch {
case errors.As(err, &syntaxErr):
read = syntaxErr.Offset
case errors.As(err, &typeErr):
read = typeErr.Offset
default:
return ""
}
before := data[:max(min(read, int64(len(data)))-1, 0)]
line := bytes.Count(before, []byte("\n")) + 1
column := len(before) - bytes.LastIndexByte(before, '\n')
return fmt.Sprintf(", line %d, column %d", line, column)
}
// write writes data to the file name in dir so that a crash at any
// moment leaves either the old file or the new one, whole: data goes to a
// temporary file in the same directory, which is synced and renamed over
// name. syncDirectory must follow, so that the rename lasts.
func write(dir, name string, data []byte) error {
path := filepath.Join(dir, name)
temporary := path + ".tmp"
err := writeSynced(temporary, data)
if err == nil {
err = os.Rename(temporary, path)
}
if err != nil {
_ = os.Remove(temporary)
}
return err
}
// syncDirectory syncs dir to the disk, so that a rename in it lasts.
func syncDirectory(dir string) error {
directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
if err != nil {
return err
}
return errors.Join(directory.Sync(), directory.Close())
}
// writeSynced writes data to the file at path, and syncs it to the disk.
func writeSynced(path string, data []byte) error {
//nolint:gosec // a state file's temporary file, in SWWAF_STATE_DIR
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
if err != nil {
return err
}
_, err = file.Write(data)
if err == nil {
err = file.Sync()
}
return errors.Join(err, file.Close())
}
File diff suppressed because it is too large Load Diff
+35
View File
@@ -0,0 +1,35 @@
package state
import (
"os"
"path/filepath"
"testing"
)
// The test is on write itself: a state file is read before it is
// written, and a directory in its place fails that read first.
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
// A directory named bans.json cannot be renamed over.
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
err = write(dir, bansJSON, []byte("{}\n"))
if err == nil {
t.Error("writing over a directory did not fail")
}
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatalf("read %s: %v", dir, err)
}
if len(entries) != 1 || entries[0].Name() != bansJSON {
t.Errorf("%s holds %v, want only bans.json", dir, entries)
}
}
+68 -13
View File
@@ -1,12 +1,15 @@
#!/bin/sh #!/bin/sh
# script/example-app: build the image, and on it the example app in # script/example-app: build the image, and on it the example app in
# deploy/example-app, then run the app's container and check that the # deploy/example-app, then run the app's container with a volume for the
# health check passes, that a request is served through smallwebwaf, # state files and check that the health check passes, that a request is
# that `sv stop` stops smallwebwaf in order, and that `docker stop` # served through smallwebwaf, that a second one in a minute bans the
# stops the container without having to kill it. The container and both # client, that a probe for /.env bans another client, which its next
# images are removed however the script ends. Building the app needs # request bans for good, that `sv stop` stops smallwebwaf in order, that
# network access, for nixpkgs' binary cache. script/check does not run # `docker stop` stops the container without having to kill it, and that
# this. # a new container on the same volume still refuses the banned client. The
# containers, the volume and both images are removed however the script
# ends. Building the app needs network access, for nixpkgs' binary cache.
# script/check does not run this.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -18,9 +21,11 @@ NAME="$("$SCRIPT_DIR/projectname")-example-$$"
IMAGE="$NAME-base" IMAGE="$NAME-base"
APP_IMAGE="$NAME-app" APP_IMAGE="$NAME-app"
CONTAINER="$NAME" CONTAINER="$NAME"
VOLUME="$NAME-state"
cleanup() { cleanup() {
docker rm --force "$CONTAINER" >/dev/null 2>&1 || true docker rm --force "$CONTAINER" >/dev/null 2>&1 || true
docker volume rm --force "$VOLUME" >/dev/null 2>&1 || true
docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true
} }
@@ -48,9 +53,42 @@ healthy() {
[ "$status" = healthy ] [ "$status" = healthy ]
} }
# logged <text>: the container's output holds text. # logged <text>...: a line of the container's output holds every text,
# in any order.
logged() { logged() {
docker logs "$CONTAINER" 2>&1 | grep -qF "$1" lines="$(docker logs "$CONTAINER" 2>&1)"
for text in "$@"; do
lines="$(printf '%s\n' "$lines" | grep -F "$text")" || return 1
done
}
# start_container: run the app's container, with the state files on the
# volume and a rate limit of one request a minute, and wait until it is
# healthy.
start_container() {
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
--volume "$VOLUME:/var/lib/smallwebwaf" \
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
"$APP_IMAGE" >/dev/null
wait_for "the health check did not pass" healthy
address="$(docker port "$CONTAINER" 8080/tcp)"
}
# refused: a request to the container gets 403, SWWAF_BAN_RESPONSE's
# default.
refused() {
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
--max-time 10 "http://$address/")" || true
[ "$code" = 403 ]
}
# refused_from <client> <path>: a request for path from client, as
# X-Forwarded-For names it, gets 403. smallwebwaf believes the header
# from docker's gateway, a private address.
refused_from() {
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
--max-time 10 --header "X-Forwarded-For: $1" "http://$address$2")" || true
[ "$code" = 403 ]
} }
main() { main() {
@@ -62,18 +100,28 @@ main() {
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \ docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
-t "$APP_IMAGE" deploy/example-app -t "$APP_IMAGE" deploy/example-app
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \ docker volume create "$VOLUME" >/dev/null
"$APP_IMAGE" >/dev/null start_container
wait_for "the health check did not pass" healthy
echo "example-app: the health check passes" echo "example-app: the health check passes"
address="$(docker port "$CONTAINER" 8080/tcp)"
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" || page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
fail "no answer on port 8080" fail "no answer on port 8080"
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page" [ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"' wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"'
echo "example-app: smallwebwaf passes a request to the app and its answer back" echo "example-app: smallwebwaf passes a request to the app and its answer back"
refused || fail "a second request in a minute was not refused"
wait_for "smallwebwaf logged no ban" logged '"action":"rate_limited"'
echo "example-app: a second request in a minute bans the client"
refused_from 203.0.113.9 /.env || fail "a probe for /.env was not refused"
wait_for "smallwebwaf logged no ban for the probe" \
logged '"action":"banned"' '"rule_ids":["env-file"]'
refused_from 203.0.113.9 / || fail "the client of the probe was let through"
wait_for "the client's next request did not make its ban permanent" \
logged '"ban_expires":"permanent"'
echo "example-app: a probe for /.env bans the client, its next request for good"
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null || docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
fail "sv stop smallwebwaf failed" fail "sv stop smallwebwaf failed"
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"' wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
@@ -83,6 +131,13 @@ main() {
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")" status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
[ "$status" = 0 ] || fail "docker stop left exit status $status" [ "$status" = 0 ] || fail "docker stop left exit status $status"
echo "example-app: docker stop stops the container in order" echo "example-app: docker stop stops the container in order"
docker rm "$CONTAINER" >/dev/null
start_container
refused || fail "the new container let the banned client through"
wait_for "smallwebwaf logged no request refused under the ban" \
logged '"action":"banned"'
echo "example-app: a new container on the same volume keeps the ban"
} }
main "$@" main "$@"
+13 -1
View File
@@ -1,6 +1,9 @@
#!/bin/sh #!/bin/sh
# script/run: build bin/smallwebwaf with script/build and run it, with # script/run: build bin/smallwebwaf with script/build and run it, with
# the settings in the environment. # the settings in the environment. Unless SWWAF_STATE_DIR is set, the
# state files go in bin/state, beside the binary, and unless
# SWWAF_RULES_DIR is set, the rule files are those of share/rules.d,
# which the image ships.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -8,6 +11,15 @@ ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)"
main() { main() {
"$SCRIPT_DIR/build" "$SCRIPT_DIR/build"
if [ -z "${SWWAF_STATE_DIR+set}" ]; then
SWWAF_STATE_DIR="$ROOT/bin/state"
export SWWAF_STATE_DIR
mkdir -p "$SWWAF_STATE_DIR"
fi
if [ -z "${SWWAF_RULES_DIR+set}" ]; then
SWWAF_RULES_DIR="$ROOT/share/rules.d"
export SWWAF_RULES_DIR
fi
exec "$ROOT/bin/smallwebwaf" exec "$ROOT/bin/smallwebwaf"
} }
+15
View File
@@ -0,0 +1,15 @@
# 00-default.rules: probes no real visitor sends, anchored at the site root
# id target action regex
env-file path ban (?i)^/\.env(\.[a-z]+)?$
vcs-dir path ban (?i)^/\.(git|svn|hg|bzr)(/|$)
secrets-dir path ban (?i)^/\.(aws|ssh|docker|kube)/
secret-file path ban (?i)^/\.(htpasswd|htaccess|npmrc|netrc|pgpass|git-credentials|bash_history|DS_Store)$
editor-dir path ban (?i)^/\.(vscode|idea)/
backup-file path ban (?i)^/[^/]+\.(php(\.[a-z0-9]+|~)|sql(\.[a-z0-9]+)?)$
log-file path ban (?i)^/(debug|error|access)\.log$
compose-file path ban (?i)^/(docker-)?compose\.ya?ml$
php-shell path ban (?i)^/(shell|c99|r57|wso|alfa)\.php$
scanner-agent user_agent ban (?i)\b(sqlmap|nikto|nuclei|masscan|zgrab|wpscan)\b
path-traversal uri block (\.\./){2,}
empty-agent user_agent log ^$
+6 -2
View File
@@ -2,10 +2,14 @@
set -euo pipefail set -euo pipefail
# runit's run script for smallwebwaf, run again whenever smallwebwaf # runit's run script for smallwebwaf, run again whenever smallwebwaf
# exits; the wait spaces out the restarts. exec, so that the signal # exits; the wait spaces out the restarts. The state directory and every
# `sv stop` sends reaches smallwebwaf itself. # file in it are given to the smallwebwaf user, so that a volume mounted
# there needs no change of owner; chown -R changes a symbolic link itself,
# never what it points to. exec, so that the signal `sv stop` sends
# reaches smallwebwaf itself.
main() { main() {
sleep 1 sleep 1
chown -R smallwebwaf:smallwebwaf "${SWWAF_STATE_DIR:-/var/lib/smallwebwaf}"
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
} }