Compare commits

...
2 Commits
Author SHA1 Message Date
clawbot 34ebf1abb9 Take in an admin's edits of the state files while running (closes #68)
check / check (push) Successful in 2m59s
smallwebwaf watches SWWAF_STATE_DIR with fsnotify and takes in a saved
edit of a state file in place of what it held. It tells its own writes
from an admin's 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 file's next write. Every ban on a netblock
is checked, and the next ban is worked out from the one that ended
last. Two metrics count the edits taken in and set aside. 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 09:48:06 +00: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
27 changed files with 2520 additions and 236 deletions
+136 -35
View File
@@ -13,19 +13,21 @@ JSON log line for every request.
Status: the first two milestones are built Status: the first two milestones are built
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and (https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are four parts of https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are five parts of
milestone 3: the static lists, the bans that broken rate limits lead to and the milestone 3: the static lists, the bans that broken rate limits lead to and the
JSON state files, which come next in the build order, and the header size and JSON state files with your edits taken in while it runs, which come next in the
the idle time as settings, which come last in it. `smallwebwaf` passes each build order, and the metrics endpoint and the header size and the idle time as
request to the app and the app's answer back, unchanged, within its timeouts and settings, which come last in it. `smallwebwaf` passes each request to the app
size limits, works out each client's address, bans a client that sends too many and the app's answer back, unchanged, within its timeouts and size limits, works
requests, refuses a client that comes from a country you refuse or from a out each client's address, bans a client that sends too many requests, refuses a
network you refuse, lets the networks you choose through, keeps its bans, each client that comes from a country you refuse or from a network you refuse, lets
client's counters and history, and GeoJS's answers in JSON files across the networks you choose through, keeps its bans, each client's counters and
restarts, and writes a JSON log line for every request. It comes as the image history, and GeoJS's answers in JSON files across restarts, takes in your edits
the app's own image is built on. The rest of the design comes after that, in the of those files while it runs, writes a JSON log line for every request, and
order of the build order in [`SPEC.md`](SPEC.md). The survey of existing tools serves Prometheus metrics to a scraper that holds the metrics token. It comes as
that led to the design is in [`EVALUATION.md`](EVALUATION.md). the image the app's own image is built on. The rest of the design comes after
that, in the order of the build order in [`SPEC.md`](SPEC.md). The survey of
existing tools that led to the design is in [`EVALUATION.md`](EVALUATION.md).
## Getting started ## Getting started
@@ -97,9 +99,9 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
seen, how many of them the ban has refused, and how many bans the netblock had seen, how many of them the ban has refused, and how many bans the netblock had
before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent; before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent;
past that, the earliest ban of the netblock that has gone longest without a past that, the earliest ban of the netblock that has gone longest without a
request is dropped first. `bans.json` shows the bans and their notes, and a request is dropped first. `bans.json` shows the bans and their notes, a
restart lifts none (see "State files" below); lifting a ban by editing it restart lifts none, and you add or lift a ban by editing it (see "State files"
comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. below).
- Refuses a request from a country you refuse with `SWWAF_BAN_RESPONSE`, as soon - Refuses a request from a country you refuse with `SWWAF_BAN_RESPONSE`, as soon
as the client's country is known and before its body is read; such a request as the client's country is known and before its body is read; such a request
is not counted for the rate limits. While one of the country lists below is is not counted for the rate limits. While one of the country lists below is
@@ -119,6 +121,14 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
limits; the country lists and bans still apply to it. limits; the country lists and bans still apply to it.
- Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any - Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any
check and without asking the app, for the image's health check. check and without asking the app, for the image's health check.
- Answers `GET /_smallwebwaf/metrics` with its metrics (see "Metrics" below) for
a request that carries `SWWAF_METRICS_TOKEN` as
`Authorization: Bearer <token>`, and with `401` for one that does not. While
the token is unset the metrics answer `404`, as does any other request under
`/_smallwebwaf/`. Unlike the health check, such a request goes through every
check any other request goes through, and is answered where another would be
passed to the app: a banned client stays refused, and each counts toward the
client's rate limits. None of them reaches the app.
- Writes a line in the request log for each request (see "Request log" below). - Writes a line in the request log for each request (see "Request log" below).
## Settings ## Settings
@@ -194,6 +204,12 @@ it, and the effective settings are logged at start.
`bans.json` is written, with every ban made in between. `bans.json` is written, with every ban made in between.
- `SWWAF_STATE_COUNTER_INTERVAL` (default `15m`): how often every state file is - `SWWAF_STATE_COUNTER_INTERVAL` (default `15m`): how often every state file is
written. written.
- `SWWAF_METRICS_TOKEN` (default unset): the token a scraper sends for the
metrics, a long random value. While it is unset the metrics are off; one
shorter than 32 characters stops the start. The settings logged at start show
`********` in its place.
- `SWWAF_METRICS_TOP_N` (default `50`): how many countries get series of their
own in the metrics by country; the others are counted as `other`.
Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024 bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
@@ -203,8 +219,8 @@ ISO 3166-1 assigns today, and `xk` for Kosovo, in either case (`de` and `DE` are
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
`su`, stops the start, and so does a code on both country lists. `off` switches `su`, stops the start, and so does a code on both country lists. `off` switches
a timeout, a size limit or a rate limit off; a timeout, a size limit or a rate limit off;
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings and the state settings `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings, the state settings
cannot be off. and `SWWAF_METRICS_TOP_N` cannot be off.
Several limits are fixed rather than settings. At most 20,000 clients are kept, Several limits are fixed rather than settings. At most 20,000 clients are kept,
with their counters and history, and an IPv6 client is counted by its /64. A new with their counters and history, and an IPv6 client is counted by its /64. A new
@@ -268,10 +284,11 @@ entries by client address, with times in UTC.
`expires` is `null`. `expires` is `null`.
- `clients.json`: each client's two buckets in the minute, the hour and the day, - `clients.json`: each client's two buckets in the minute, the hour and the day,
and its history: when it was first and last seen, its country as last looked and its history: when it was first and last seen, its country as last looked
up and when, its requests, how many were forwarded and how many refused, the up and when, its requests, how many were forwarded and how many refused (one
body bytes in each direction, its responses by status class and its offences `smallwebwaf` answered at its own endpoints is neither, unless it was refused
by kind. Each client is on a line of its own, so `grep` shows everything about with `401` for a missing or wrong token), the body bytes in each direction,
one. its responses by status class and its offences by kind. Each client is on a
line of its own, so `grep` shows everything about one.
- `lookups.json`: GeoJS's answers, one to a line, with when GeoJS gave each and - `lookups.json`: GeoJS's answers, one to a line, with when GeoJS gave each and
when it was last used. when it was last used.
@@ -295,9 +312,89 @@ without a field it needs, named with the entry's place in the file: a ban's
`netblock`, `start` or `expires`, which is `null` for a permanent ban; a `netblock`, `start` or `expires`, which is `null` for a permanent ban; a
client's `client`, or the `start` of a window in which it has requests; an client's `client`, or the `start` of a window in which it has requests; an
answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or
`answered`. An edit made while `smallwebwaf` runs is overwritten by its next `answered`. The AS number and AS name come with their lookup.
write: taking it in comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
The AS number and AS name come with their lookup. While it runs, `smallwebwaf` watches `SWWAF_STATE_DIR` and takes in your edit of
a state file as soon as you save it: what the file then holds replaces what
`smallwebwaf` held for it, as if read at start. It tells its own writes from
yours by comparing the file with what it last read or wrote, and before it
writes a file it takes in any edit made since, so your edit is not overwritten;
a change `smallwebwaf` made after you opened the file, such as a new ban, is
lost when you save over it. An edit that would stop the start, because it does
not parse, has another `version` or leaves out a field an entry needs, does not
stop the running `smallwebwaf`: it keeps what it holds, and at the file's next
write renames your file to `<name>.bad`, such as `bans.json.bad`, writes the
file again from memory, and logs the file and where the error is. It waits for
that write because an editor's file can be read before the editor has finished
writing it. Mend the `.bad` file and move it back. A file you remove is written
again at its next write.
To ban a netblock, add an entry to `bans.json` with its `netblock`, its `start`
and its `expires`, `null` for a ban that never ends; its `notes` may be left
out. This `bans.json` bans `203.0.113.0/24` for good:
```json
{
"version": 1,
"bans": [
{
"netblock": "203.0.113.0/24",
"start": "2026-10-06T12:00:00Z",
"expires": null
}
]
}
```
To lift a ban, delete its entry. `smallwebwaf` then forgets the ban, so it does
not make the netblock's next ban longer.
## Metrics
`GET /_smallwebwaf/metrics` answers with the metrics in the Prometheus text
format, for a scraper that sends `SWWAF_METRICS_TOKEN`, through traefik like any
other request. No metric carries a client's address.
- `smallwebwaf_requests_total`, `smallwebwaf_request_bytes_total` and
`smallwebwaf_response_bytes_total`: requests, and their body bytes each way,
by `status_class`, such as `2xx`, or `none` when nothing was sent, and by
`action`, as the request log names it.
- `smallwebwaf_request_duration_seconds`: how long requests took, and
`smallwebwaf_upstream_duration_seconds`: how long those passed to the app took
from then on, as histograms; `smallwebwaf_requests_in_flight`: the requests
under way.
- `smallwebwaf_rate_limit_hits_total` by `window`,
`smallwebwaf_size_and_time_limit_hits_total` by `limit`, the setting whose
limit was passed, `smallwebwaf_offences_total` by `kind`, and
`smallwebwaf_bans_made_total` by `cause`; `smallwebwaf_active_bans` and
`smallwebwaf_permanent_bans`.
- `smallwebwaf_country_requests_total`,
`smallwebwaf_country_request_bytes_total`,
`smallwebwaf_country_response_bytes_total`, and
`smallwebwaf_country_list_refusals_total`, the requests the country lists
refused, by `country`, for the requests whose client's country is known. The
`SWWAF_METRICS_TOP_N` countries with the most requests since the start have
series of their own, and the others are counted as `other`. A country that
drops out of them loses its series, and its later requests count as `other`;
one that comes into them gets a series that counts from then on.
- `smallwebwaf_geojs_requests_total`: the requests to GeoJS;
`smallwebwaf_geojs_failures_total`: those that failed, an answer that leaves
out an address asked about included; and `smallwebwaf_geojs_unanswered_total`:
the requests whose client counted as coming from an unknown country because
GeoJS had not answered about it in time.
- `smallwebwaf_tracked_clients`: the clients in the table of clients.
- `smallwebwaf_state_file_writes_total`,
`smallwebwaf_state_file_write_failures_total`,
`smallwebwaf_state_file_last_write_timestamp_seconds` and
`smallwebwaf_state_file_size_bytes`, by `file`; and, by `file` too,
`smallwebwaf_state_file_edits_taken_in_total`: your edits taken in, and
`smallwebwaf_state_file_edits_set_aside_total`: those renamed to `<name>.bad`
because they would stop the start.
- Go's own `go_` metrics and the process's `process_` metrics.
The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request
log") are not counted. The metrics of the features still to come, such as the
rule files, come with them.
## Why ## Why
@@ -400,9 +497,9 @@ goes through the candidates one by one.
readable JSON files, written regularly and at every stop, so a restart loses readable JSON files, written regularly and at every stop, so a restart loses
nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks
up the change. Nothing is read from disk while serving a request. The files up the change. Nothing is read from disk while serving a request. The files
for the bans, the clients and the GeoJS answers are built (see "State files" for the bans, the clients and the GeoJS answers are built, with an edit taken
above); the others come with their features, and taking in an edit while in while running (see "State files" above); the others come with their
running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. features.
- Health checks, the metrics, and listing, adding and lifting bans or asking why - Health checks, the metrics, and listing, adding and lifting bans or asking why
a given address was refused, all on the one port every request uses: under a given address was refused, all on the one port every request uses: under
`/_smallwebwaf/` on the app's own address, through traefik like any other `/_smallwebwaf/` on the app's own address, through traefik like any other
@@ -583,15 +680,18 @@ addresses are never sent to GeoJS.
limits, and writes the request's log line. Its `check` method is where a limits, and writes the request's log line. Its `check` method is where a
request is refused before anything reaches the app: for `SWWAF_DENY_NETS`, for request is refused before anything reaches the app: for `SWWAF_DENY_NETS`, for
a ban, for the country lists, for a rate limit, which bans the client, and for a ban, for the country lists, for a rate limit, which bans the client, and for
an announced body over the size limit. an announced body over the size limit. A request under `/_smallwebwaf/` that
`check` lets through is answered by `answerAdmin` instead of reaching the app.
- `internal/metrics`: the metrics, counted as the other parts tell it what
happened, and served in the Prometheus text format.
- `internal/bans`: the ban ledger: each netblock's bans with their notes, how - `internal/bans`: the ban ledger: each netblock's bans with their notes, how
long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held. long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held.
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the - `internal/lookup`: looks up each client's country through GeoJS, and keeps the
answers. answers.
- `internal/ratelimit`: the table of clients: counts each client's requests, - `internal/ratelimit`: the table of clients: counts each client's requests,
tells when one takes it over a rate limit, and keeps each client's history. tells when one takes it over a rate limit, and keeps each client's history.
- `internal/state`: reads the state files at start, and writes them when they - `internal/state`: reads the state files at start, takes in an admin's edit of
are due and at the stop. one while running, and writes them when they are due and at the stop.
- `internal/requestlog`: the lines on stdout: the request log line and the - `internal/requestlog`: the lines on stdout: the request log line and the
process's own messages. process's own messages.
- `Dockerfile`: the lint and test phases, then the image, whose last stage - `Dockerfile`: the lint and test phases, then the image, whose last stage
@@ -602,8 +702,10 @@ addresses are never sent to GeoJS.
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
table of clients to 20,000, the GeoJS answers to 100,000 and the banned table of clients to 20,000, the GeoJS answers to 100,000 and the banned
netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen. The country netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and
codes are the list in `internal/config/config.go`. `github.com/prometheus/client_golang` keeps the metrics and serves them, and
`github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file is saved.
The country codes are the list in `internal/config/config.go`.
## Entrypoints ## Entrypoints
@@ -646,9 +748,8 @@ so that they run in minimal containers.
## TODO ## TODO
- The rest of milestone 3, from taking in an admin's edits to the state files - The rest of milestone 3, from exemptions up to the rest of the request log's
(https://git.eeqj.de/sneak/smallwebwaf/issues/68) up to the metrics endpoint, fields, and the rest of the design, in the order of the build order in
and the rest of the design, in the order of the build order in
[`SPEC.md`](SPEC.md). [`SPEC.md`](SPEC.md).
## Documents ## Documents
+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
)
+39
View File
@@ -1,2 +1,41 @@
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.13.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
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=
+87 -26
View File
@@ -26,9 +26,9 @@ const maxTextBytes = 256
type Rules struct { type Rules struct {
// LimitBanDuration is how long a first ban lasts. // LimitBanDuration is how long a first ban 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 a broken limit counts as a repeat, which bans
// repeatFactor times as long as that ban. // 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; a ban that would be longer is
// permanent instead. // permanent instead.
@@ -112,6 +112,8 @@ type Ledger struct {
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, at most rules.MaxBans.
held int held int
// made is how many bans BanForLimit has made since the start.
made int
// v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6 // v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6
// netblocks that have been banned. Check looks for a ban at each of // 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 // them, so that a ban read from bans.json refuses every client in its
@@ -160,23 +162,35 @@ func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
continue continue
} }
// A ban is made only once the one before has ended, so only the ban := activeBan(*bans, now)
// last can be active. if ban != nil {
last := &(*bans)[len(*bans)-1] ban.Notes.Requests++
if last.ActiveAt(now) { ban.Notes.Refused++
last.Notes.Requests++
last.Notes.Refused++
return *last, true return *ban, true
} }
} }
return Ban{}, false return Ban{}, false
} }
// 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. A first ban lasts LimitBanDuration. A ban made within
// LimitBanRepeatWindow after the netblock's last ban ended lasts // LimitBanRepeatWindow after the netblock's ban that ended last 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
@@ -190,12 +204,22 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
bans, found := l.netblocks.Get(netblock) bans, found := l.netblocks.Get(netblock)
if found { if found {
last = &(*bans)[len(*bans)-1] active := activeBan(*bans, now)
if last.ActiveAt(now) { if active != nil {
return *last return *active
} }
notes.EarlierBans = last.Notes.EarlierBans + 1 // No ban is active, so each has an end. A ban an admin adds to
// bans.json can start after another and end before it, so the
// ban that ended last is looked for among them all.
ended := slices.MaxFunc(*bans, func(a, b Ban) int {
return a.Expires.Compare(b.Expires)
})
last = &ended
// The netblock's first ban held counts the bans it had before that
// one, since dropped to make room, and each ban held adds one.
notes.EarlierBans = (*bans)[0].Notes.EarlierBans + len(*bans)
} }
notes.Request = notes.Request.cut() notes.Request = notes.Request.cut()
@@ -206,6 +230,7 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
Notes: notes, Notes: notes,
} }
l.add(ban) l.add(ban)
l.made++
select { select {
case l.changed <- struct{}{}: case l.changed <- struct{}{}:
@@ -229,6 +254,38 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans) return slices.Clone(*bans)
} }
// Made returns how many bans the ledger has made since the start; bans
// read from bans.json are not among them.
func (l *Ledger) Made() int {
l.mu.Lock()
defer l.mu.Unlock()
return l.made
}
// Count returns how many of the bans held are active at now, and how many
// are permanent.
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) {
active++
}
if ban.Permanent() {
permanent++
}
}
}
return active, permanent
}
// Snapshot returns every ban held, sorted by netblock, and each // Snapshot returns every ban held, sorted by netblock, and each
// netblock's bans oldest first, as bans.json lists them. // netblock's bans oldest first, as bans.json lists them.
func (l *Ledger) Snapshot() []Ban { func (l *Ledger) Snapshot() []Ban {
@@ -247,21 +304,25 @@ func (l *Ledger) Snapshot() []Ban {
return held return held
} }
// Load puts bans read from bans.json into a ledger that holds none yet, // Load puts bans read from bans.json into the ledger, in place of the
// in the order they started, so that a netblock whose last ban started // bans it holds, in the order they started, so that a netblock whose last
// latest counts as the most recently seen. Each netblock is masked to its // ban started latest counts as the most recently seen. Each netblock is
// length, so that 203.0.113.9/24 is 203.0.113.0/24, and each text in the // masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24, and
// notes is cut to 256 bytes. Past MaxBans the earliest bans are dropped, // each text in the notes is cut to 256 bytes. Past MaxBans the earliest
// as when they are made. // bans are dropped, as when they are made.
func (l *Ledger) Load(bans []Ban) { func (l *Ledger) Load(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
bans = slices.Clone(bans) bans = slices.Clone(bans)
slices.SortStableFunc(bans, func(a, b Ban) int { slices.SortStableFunc(bans, func(a, b Ban) int {
return a.Start.Compare(b.Start) return a.Start.Compare(b.Start)
}) })
l.mu.Lock()
defer l.mu.Unlock()
l.netblocks.Purge()
l.held = 0
l.v4Lengths, l.v6Lengths = nil, nil
for _, ban := range bans { for _, ban := range bans {
ban.Netblock = ban.Netblock.Masked() ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut() ban.Notes.Request = ban.Notes.Request.cut()
@@ -298,8 +359,8 @@ func (l *Ledger) add(ban Ban) {
} }
// expiry returns when a ban for a broken limit made at now ends, or zero // expiry returns when a ban for a broken limit made at now ends, or zero
// when it is permanent. last is the netblock's last ban, which has ended, // when it is permanent. last is the netblock's ban that ended last, or nil
// or nil when it has none. // when it has none.
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time { func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration length := l.rules.LimitBanDuration
+98
View File
@@ -130,6 +130,69 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
} }
} }
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)
ban, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), 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 notes.
netblock := netip.MustParsePrefix("203.0.113.9/32")
nineHours := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(9 * time.Hour),
Notes: bans.Notes{EarlierBans: 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, it, 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 != 4 {
t.Errorf("the next ban lasts %s with %d earlier bans, want 27h and 4",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
}
func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) { func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
t.Parallel() t.Parallel()
@@ -151,6 +214,41 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
} }
} }
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()}
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.0/24"), Start: midnight()},
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) { func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
t.Parallel() t.Parallel()
+34
View File
@@ -17,6 +17,7 @@ import (
"strconv" "strconv"
"strings" "strings"
"time" "time"
"unicode/utf8"
) )
// 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
@@ -103,6 +104,12 @@ type Config struct {
StateDir string StateDir string
StateWriteDelay time.Duration StateWriteDelay time.Duration
StateCounterInterval time.Duration StateCounterInterval time.Duration
// 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
// settings are the values read, as given or by default, for the // settings are the values read, as given or by default, for the
// log line at start. // log line at start.
@@ -119,6 +126,10 @@ 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.
masked = "********"
) )
var ( var (
@@ -150,6 +161,7 @@ var (
"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( errNotAbsolutePath = errors.New(
"is not an absolute path, such as /var/lib/smallwebwaf") "is not an absolute path, such as /var/lib/smallwebwaf")
errShortToken = errors.New("is shorter than 32 characters")
) )
// FromEnvironment reads the settings with lookupEnv, normally // FromEnvironment reads the settings with lookupEnv, normally
@@ -188,6 +200,8 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"), StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"), StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"), StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
} }
for _, country := range cfg.ExclusivelyAllowedCountries { for _, country := range cfg.ExclusivelyAllowedCountries {
@@ -353,6 +367,26 @@ func (e *environment) absolutePath(name, defaultValue string) string {
return path 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.lookupEnv(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
}
// 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) {
+54 -1
View File
@@ -44,8 +44,13 @@ const (
stateDir = "SWWAF_STATE_DIR" stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
) )
// token is a token of 32 characters, the shortest allowed.
const token = "0123456789abcdef0123456789abcdef"
// off switches a timeout, a size limit or a rate limit off. // off switches a timeout, a size limit or a rate limit off.
const off = "off" const off = "off"
@@ -98,6 +103,8 @@ func TestDefaults(t *testing.T) {
StateDir: "/var/lib/smallwebwaf", StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second, StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute, StateCounterInterval: 15 * time.Minute,
MetricsToken: "",
MetricsTopN: 50,
}) })
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" { if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -145,6 +152,8 @@ func TestValuesAsSet(t *testing.T) {
stateDir: "/srv/waf-state", stateDir: "/srv/waf-state",
stateWriteDelay: "500ms", stateWriteDelay: "500ms",
stateCounterInterval: "1h", stateCounterInterval: "1h",
metricsToken: token,
metricsTopN: "10",
}) })
wantSettings(t, cfg, config.Config{ wantSettings(t, cfg, config.Config{
@@ -169,6 +178,8 @@ func TestValuesAsSet(t *testing.T) {
StateDir: "/srv/waf-state", StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond, StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour, StateCounterInterval: time.Hour,
MetricsToken: token,
MetricsTopN: 10,
}) })
if cfg.UpstreamURL.String() != "https://app.internal:8443/" { if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -344,6 +355,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"}, {stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
{stateWriteDelay, off}, {stateWriteDelay, "0s"}, {stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"}, {stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
} { } {
t.Run(tc.name+"="+tc.value, func(t *testing.T) { t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -360,6 +372,39 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
} }
} }
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
// Characters are counted, not bytes: each é takes two.
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv)
want := metricsToken + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestTokenIsLoggedMasked(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{metricsToken: token})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
if strings.Contains(out.String(), token) ||
!strings.Contains(out.String(), `"`+metricsToken+`":"********"`) {
t.Errorf("the token is not logged masked: %s", out.String())
}
}
func TestLogsEachSettingWithItsValue(t *testing.T) { func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Parallel() t.Parallel()
@@ -407,6 +452,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
stateDir: "/var/lib/smallwebwaf", stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s", stateWriteDelay: "10s",
stateCounterInterval: "15m", stateCounterInterval: "15m",
metricsToken: "",
metricsTopN: "50",
} }
if !maps.Equal(line.Settings, want) { if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want) t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -435,7 +482,8 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
wantBanSettings(t, got, want) wantBanSettings(t, got, want)
} }
// wantBanSettings checks the settings for bans and the state files. // wantBanSettings checks the settings for bans, the state files and the
// metrics.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) { func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper() t.Helper()
@@ -453,6 +501,11 @@ func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
got.StateCounterInterval != want.StateCounterInterval { got.StateCounterInterval != want.StateCounterInterval {
t.Errorf("state settings\n%+v\nwant\n%+v", got, want) t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
} }
if got.MetricsToken != want.MetricsToken || got.MetricsTopN != want.MetricsTopN {
t.Errorf("metrics token %q and top %d, want %q and %d",
got.MetricsToken, got.MetricsTopN, want.MetricsToken, want.MetricsTopN)
}
} }
// wantNetblocks checks a list of netblocks. // wantNetblocks checks a list of netblocks.
+24 -5
View File
@@ -19,6 +19,7 @@ import (
"time" "time"
"github.com/hashicorp/golang-lru/v2/simplelru" "github.com/hashicorp/golang-lru/v2/simplelru"
"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,
@@ -64,6 +65,9 @@ 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
} }
// GeoJS looks up clients' countries through GeoJS. At most one request // GeoJS looks up clients' countries through GeoJS. At most one request
@@ -73,6 +77,7 @@ type GeoJS struct {
url string url string
now func() time.Time now func() time.Time
processLog *slog.Logger processLog *slog.Logger
metrics *metrics.Metrics
// 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
@@ -121,6 +126,7 @@ 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,
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
@@ -160,6 +166,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 {
@@ -188,19 +197,21 @@ func (g *GeoJS) Snapshot() []Answer {
return answers return answers
} }
// Load keeps answers read from lookups.json, in a GeoJS that keeps none // Load keeps answers read from lookups.json, in place of the answers it
// yet, in the order they were last used, so that the one used longest // 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 // ago is dropped first. Answers GeoJS gave keepFor ago or more are
// dropped. // dropped.
func (g *GeoJS) Load(answers []Answer) { func (g *GeoJS) Load(answers []Answer) {
g.mu.Lock()
defer g.mu.Unlock()
answers = slices.Clone(answers) answers = slices.Clone(answers)
slices.SortStableFunc(answers, func(a, b Answer) int { slices.SortStableFunc(answers, func(a, b Answer) int {
return a.Used.Compare(b.Used) return a.Used.Compare(b.Used)
}) })
g.mu.Lock()
defer g.mu.Unlock()
g.answers.Purge()
now := g.now() now := g.now()
for _, answer := range answers { for _, answer := range answers {
@@ -234,6 +245,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
} }
@@ -243,6 +256,8 @@ func (g *GeoJS) answerOrWait(
} }
if w.late { if w.late {
g.metrics.GeoJSUnanswered.Inc()
return "", nil return "", nil
} }
@@ -355,6 +370,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)
@@ -399,6 +416,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
+49
View File
@@ -13,7 +13,9 @@ import (
"testing/synctest" "testing/synctest"
"time" "time"
"github.com/prometheus/client_golang/prometheus/testutil"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
) )
const ( const (
@@ -194,6 +196,7 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
URL: lookup.URL, URL: lookup.URL,
Now: time.Now, Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1),
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -348,6 +351,40 @@ func TestAtMost10000ClientsWait(t *testing.T) {
}) })
} }
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,
})
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.
const ( const (
answering = iota answering = iota
@@ -493,6 +530,7 @@ func start() (*standIn, *testClock, *lookup.GeoJS) {
URL: lookup.URL, URL: lookup.URL,
Now: clock.Now, Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler), ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1),
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -549,6 +587,17 @@ func wantAsked(t *testing.T, geojs *standIn, i int, want ...string) {
} }
} }
// wantUnanswered checks how many requests m counts as having gone without
// 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, // waitForRequests waits until g has done all it can before time passes,
// checks that GeoJS has had count requests, and returns the addresses each // checks that GeoJS has had count requests, and returns the addresses each
// asked about. // asked about.
+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
}
+277
View File
@@ -0,0 +1,277 @@
// 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/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// 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
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,
// 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,
) {
m.registry.MustRegister(
// Every ban smallwebwaf makes so far is for a broken limit.
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_bans_made_total",
Help: "Bans made, by cause.",
ConstLabels: prometheus.Labels{"cause": "limit"},
}, func() float64 {
return float64(ledger.Made())
}),
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.",
}, 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())
}),
)
}
// 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)
}
}
// 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)
}
+43
View File
@@ -0,0 +1,43 @@
package proxy
import (
"crypto/subtle"
"net/http"
"strings"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// answerAdmin answers a request for smallwebwaf itself, under
// /_smallwebwaf/, once it has passed the checks: GET MetricsPath with
// SWWAF_METRICS_TOKEN gets the metrics, and without it is refused with
// 401. Any other request gets 404, as the metrics do while
// SWWAF_METRICS_TOKEN is unset.
func (rq *request) answerAdmin() {
rq.line.Action = requestlog.ActionAdmin
rq.startClientResponseTimeout()
token := rq.h.config.MetricsToken
switch {
case token == "" || rq.in.Method != http.MethodGet || rq.in.URL.Path != MetricsPath:
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:
rq.h.metrics.ServeHTTP(rq.out, rq.in)
}
}
// 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
}
+25 -7
View File
@@ -402,36 +402,54 @@ 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()
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, "GET "+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")
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, string(got.body)
} }
+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
+21
View File
@@ -86,6 +86,27 @@ func TestHealthEndpointIsNotInTheHistory(t *testing.T) {
} }
} }
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. // historyOf returns the history of the client at addr.
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History { func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
t.Helper() t.Helper()
+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)
}) })
} }
} }
+455
View File
@@ -0,0 +1,455 @@
package proxy_test
import (
"io"
"net/http"
"net/http/httptest"
"net/netip"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"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 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))
}
+28 -2
View File
@@ -8,11 +8,13 @@ import (
"log" "log"
"log/slog" "log/slog"
"net/http" "net/http"
"strings"
"time" "time"
"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"
) )
@@ -23,10 +25,18 @@ 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"
// Params are what New needs. // Params are what New needs.
type Params struct { type Params struct {
Config *config.Config Config *config.Config
@@ -44,13 +54,14 @@ type Params struct {
} }
// Server is the server smallwebwaf runs, with the parts of the proxy // Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep. // whose state the state files keep, and the metrics.
type Server struct { type Server struct {
*http.Server *http.Server
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Metrics *metrics.Metrics
} }
// New returns the server smallwebwaf runs: each request it reads passes // New returns the server smallwebwaf runs: each request it reads passes
@@ -61,6 +72,7 @@ type Server struct {
// applies the timeouts and size limits from then on. // applies the timeouts and size limits from then on.
func New(params Params) *Server { func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn) errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
m := metrics.New(params.Config.MetricsTopN)
h := &handler{ h := &handler{
config: params.Config, config: params.Config,
requestLog: params.RequestLog, requestLog: params.RequestLog,
@@ -68,6 +80,7 @@ func New(params Params) *Server {
errorLog: errorLog, errorLog: errorLog,
transport: newTransport(), transport: newTransport(),
now: params.Now, now: params.Now,
metrics: m,
limiter: ratelimit.New(ratelimit.Limits{ limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute, PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour, PerHour: params.Config.RateLimitPerHour,
@@ -83,8 +96,10 @@ func New(params Params) *Server {
URL: params.GeoJSURL, URL: params.GeoJSURL,
Now: params.Now, Now: params.Now,
ProcessLog: params.ProcessLog, ProcessLog: params.ProcessLog,
Metrics: m,
}), }),
} }
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
return &Server{ return &Server{
Server: &http.Server{ Server: &http.Server{
@@ -102,6 +117,7 @@ func New(params Params) *Server {
Ledger: h.ledger, Ledger: h.ledger,
Limiter: h.limiter, Limiter: h.limiter,
GeoJS: h.geojs, GeoJS: h.geojs,
Metrics: m,
} }
} }
@@ -114,6 +130,7 @@ 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
@@ -133,7 +150,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()
@@ -157,5 +175,13 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
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())
} }
+51 -17
View File
@@ -22,11 +22,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
@@ -65,9 +67,11 @@ type request struct {
requestSent time.Time requestSent 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, and works out the client.
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
@@ -141,6 +145,7 @@ func (rq *request) check(ctx context.Context) *refusal {
return &refusal{ return &refusal{
status: http.StatusRequestEntityTooLarge, status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge, action: requestlog.ActionTooLarge,
limit: "SWWAF_REQUEST_MAX_BYTES",
} }
} }
@@ -206,7 +211,11 @@ func (rq *request) modifyResponse(res *http.Response) error {
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
} }
@@ -278,7 +287,8 @@ func (rq *request) refuse(r refusal) {
rq.cancel() 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()
@@ -295,24 +305,37 @@ func (rq *request) finish() {
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)
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 = requestlog.Milliseconds(upstreamDuration)
} }
// 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())
@@ -327,9 +350,12 @@ func (rq *request) addToHistory() {
requestBytes = rq.body.bytes.Load() requestBytes = rq.body.bytes.Load()
} }
forwarded := !rq.upstreamStart.IsZero()
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{ rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country, Country: rq.line.Country,
Forwarded: !rq.upstreamStart.IsZero(), Forwarded: forwarded,
Refused: !forwarded && rq.refused.Load() != nil,
Status: rq.out.status, Status: rq.out.status,
RequestBytes: requestBytes, RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes, ResponseBytes: rq.out.bytes,
@@ -369,21 +395,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()
@@ -395,6 +426,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
@@ -403,6 +435,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
@@ -452,6 +485,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",
}) })
} }
} }
+21 -12
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.
@@ -37,29 +39,25 @@ func TestRequestTimeouts(t *testing.T) {
}{ }{
{ {
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,6 +280,7 @@ 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) { func TestClosesAnIdleConnection(t *testing.T) {
+8 -5
View File
@@ -19,26 +19,29 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100}, {Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101}, {Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5}, {Forwarded: true, Status: 304, RequestBytes: 5},
{Country: "FR", Status: 403, ResponseBytes: 10, BrokeLimit: true}, {Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12}, {Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response. // Closed without an answer: refused, and no response.
{Status: 0}, {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) limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
} }
want := ratelimit.History{ want := ratelimit.History{
FirstSeen: start, FirstSeen: start,
LastSeen: start.Add(5 * time.Minute), LastSeen: start.Add(6 * time.Minute),
Country: "FR", Country: "FR",
LookedUp: start.Add(3 * time.Minute), LookedUp: start.Add(3 * time.Minute),
Requests: 6, Requests: 7,
Forwarded: 4, Forwarded: 4,
Refused: 2, Refused: 2,
RequestBytes: 15, RequestBytes: 15,
ResponseBytes: 122, ResponseBytes: 122,
Responses: ratelimit.Responses{ Responses: ratelimit.Responses{
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 1, Status5xx: 1, Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 2, Status5xx: 1,
}, },
Offences: ratelimit.Offences{Limit: 1}, Offences: ratelimit.Offences{Limit: 1},
} }
+28 -10
View File
@@ -71,7 +71,9 @@ type History struct {
Country string `json:"country,omitempty"` Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"` LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to // Requests are all the client's requests: Forwarded those passed to
// the app, Refused those refused before anything reached it. // 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"` Requests int64 `json:"requests"`
Forwarded int64 `json:"forwarded"` Forwarded int64 `json:"forwarded"`
Refused int64 `json:"refused"` Refused int64 `json:"refused"`
@@ -103,9 +105,12 @@ type Offences struct {
type Request struct { type Request struct {
// Country is the client's country, when the request looked it up. // Country is the client's country, when the request looked it up.
Country string Country string
// Forwarded is true for a request passed to the app, false for one // Forwarded is true for a request passed to the app, Refused for one
// refused before anything reached it. // 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 Forwarded bool
Refused bool
// 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
// RequestBytes and ResponseBytes are the body bytes of the request // RequestBytes and ResponseBytes are the body bytes of the request
@@ -199,7 +204,9 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
h.Requests++ h.Requests++
if r.Forwarded { if r.Forwarded {
h.Forwarded++ h.Forwarded++
} else { }
if r.Refused {
h.Refused++ h.Refused++
} }
@@ -235,6 +242,14 @@ func (l *Limiter) Requests(netblock netip.Prefix) int64 {
return requests return requests
} }
// 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 // Snapshot returns every client in the table, sorted by address, as
// clients.json lists them. // clients.json lists them.
func (l *Limiter) Snapshot() []Client { func (l *Limiter) Snapshot() []Client {
@@ -254,18 +269,21 @@ func (l *Limiter) Snapshot() []Client {
return clients return clients
} }
// Load puts clients read from clients.json into a table that holds none // Load puts clients read from clients.json into the table, in place of
// yet, in the order they were last seen, so that the least recently seen // the clients it holds, in the order they were last seen, so that the
// is dropped first. Buckets whose time has passed at now are emptied. // least recently seen is dropped first. Buckets whose time has passed at
// now are emptied.
func (l *Limiter) Load(clients []Client, now time.Time) { func (l *Limiter) Load(clients []Client, now time.Time) {
l.mu.Lock()
defer l.mu.Unlock()
clients = slices.Clone(clients) clients = slices.Clone(clients)
slices.SortStableFunc(clients, func(a, b Client) int { slices.SortStableFunc(clients, func(a, b Client) int {
return a.History.LastSeen.Compare(b.History.LastSeen) return a.History.LastSeen.Compare(b.History.LastSeen)
}) })
l.mu.Lock()
defer l.mu.Unlock()
l.clients.Purge()
for _, c := range clients { for _, c := range clients {
for i, b := range c.buckets() { for i, b := range c.buckets() {
// The window that ends at now covers neither bucket once it // The window that ends at now covers neither bucket once it
+19 -9
View File
@@ -89,6 +89,7 @@ func Run(ctx context.Context, params Params) int {
GeoJS: server.GeoJS, GeoJS: server.GeoJS,
Now: now, Now: now,
ProcessLog: processLog, ProcessLog: processLog,
Metrics: server.Metrics,
}) })
if err != nil { if err != nil {
processLog.Error("cannot use the state files", "error", err.Error()) processLog.Error("cannot use the state files", "error", err.Error())
@@ -112,9 +113,10 @@ func Run(ctx context.Context, params Params) int {
return serve(ctx, server.Server, listener, files, processLog) return serve(ctx, server.Server, listener, files, processLog)
} }
// serve serves requests on listener, and writes the state files as they // serve serves requests on listener, writes the state files as they are
// are due, until ctx is done. Then it gives the requests in progress // due, and takes in an admin's edits of them, until ctx is done. Then it
// shutdownTimeout to finish, and writes every state file. // gives the requests in progress shutdownTimeout to finish, and writes
// every state file.
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, processLog *slog.Logger, files *state.Files, processLog *slog.Logger,
@@ -129,12 +131,18 @@ func serve(
defer stopWriting() defer stopWriting()
written := make(chan struct{}) written := make(chan struct{})
watched := make(chan struct{})
go func() { go func() {
files.Run(writing) files.Run(writing)
close(written) close(written)
}() }()
go func() {
files.Watch(writing)
close(watched)
}()
select { select {
case err := <-served: case err := <-served:
processLog.Error("serving failed", "error", err.Error()) processLog.Error("serving failed", "error", err.Error())
@@ -164,13 +172,15 @@ func serve(
return 1 return 1
} }
// Run's last write has ended, so nothing else writes the files. Every // Run and Watch have ended, so nothing else reads or writes the
// request has ended too, but for two kinds that Go's server does not // files. Every request has ended too, but for two kinds
// wait for: one cut off because Shutdown timed out, and one whose // that Go's server does not wait for: one cut off because Shutdown
// connection switched protocols, such as a WebSocket. Such a request // timed out, and one whose connection switched protocols, such as a
// adds to its client's history only as it ends, which can be after // WebSocket. Such a request adds to its client's history only as it
// this write, and then that request is missing from clients.json. // ends, which can be after this write, and then that request is
// missing from clients.json.
<-written <-written
<-watched
err = files.WriteAll() err = files.WriteAll()
if err != nil { if err != nil {
+98 -9
View File
@@ -29,7 +29,10 @@ const (
localhost = "127.0.0.1" localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR" listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL" upstreamURL = "SWWAF_UPSTREAM_URL"
trustedProxies = "SWWAF_TRUSTED_PROXIES"
stateDir = "SWWAF_STATE_DIR" stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
// greeting is what the tests' app answers. // greeting is what the tests' app answers.
greeting = "hello from the app" greeting = "hello from the app"
@@ -121,6 +124,28 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
} }
} }
func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) {
t.Parallel()
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
out := &output{}
status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
}
line := out.line(t, "msg", "invalid setting")
if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" {
t.Errorf("start refused with %v", line)
}
if strings.Contains(out.text(), token) {
t.Errorf("the output shows the token:\n%s", out.text())
}
}
func TestAddressInUseStopsTheStart(t *testing.T) { func TestAddressInUseStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -195,8 +220,8 @@ func TestStateKeptAcrossRestarts(t *testing.T) {
rateLimitPerDay: "2", rateLimitPerDay: "2",
// Neither comes due in the test: the files are written as // Neither comes due in the test: the files are written as
// smallwebwaf stops. // smallwebwaf stops.
"SWWAF_STATE_WRITE_DELAY": "1h", stateWriteDelay: "1h",
"SWWAF_STATE_COUNTER_INTERVAL": "1h", stateCounterInterval: "1h",
} }
// The two requests a day allows, and a stop. // The two requests a day allows, and a stop.
@@ -228,7 +253,7 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
listenAddr: localhost + ":0", listenAddr: localhost + ":0",
upstreamURL: startApp(t), upstreamURL: startApp(t),
stateDir: t.TempDir(), stateDir: t.TempDir(),
"SWWAF_TRUSTED_PROXIES": localhost + "/32", trustedProxies: localhost + "/32",
rateLimitPerDay: "1", rateLimitPerDay: "1",
scope: "24", scope: "24",
} }
@@ -259,6 +284,38 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
}) })
} }
func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) {
t.Parallel()
const (
// bans.json as an admin writes it with a ban, permanent, on
// 203.0.113.0/24, and with none.
oneBan = `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", ` +
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`
noBan = `{"version": 1, "bans": []}`
)
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
trustedProxies: localhost + "/32",
// No write comes due in the test, so only the watch on the
// directory can take the edits in.
stateWriteDelay: "1h",
stateCounterInterval: "1h",
}
runUntilStopped(t, env, func(url string) {
path := filepath.Join(dir, "bans.json")
saveUntilAnswered(t, path, oneBan, url, "203.0.113.9", http.StatusForbidden)
wantStatus(t, url, "198.51.100.7", http.StatusOK)
saveUntilAnswered(t, path, noBan, url, "203.0.113.9", http.StatusOK)
})
}
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) { func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -361,9 +418,9 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
listenAddr: localhost + ":0", listenAddr: localhost + ":0",
upstreamURL: appURL, upstreamURL: appURL,
stateDir: dir, stateDir: dir,
"SWWAF_STATE_WRITE_DELAY": "10s", stateWriteDelay: "10s",
"SWWAF_STATE_COUNTER_INTERVAL": "15m", stateCounterInterval: "15m",
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K", "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
"SWWAF_CLIENT_IDLE_TIMEOUT": "120s", "SWWAF_CLIENT_IDLE_TIMEOUT": "120s",
@@ -456,6 +513,40 @@ func wantRefused(t *testing.T, url string) {
func wantStatus(t *testing.T, url, from string, status int) { func wantStatus(t *testing.T, url, from string, status int) {
t.Helper() t.Helper()
got := statusFrom(t, url, from)
if got != status {
t.Errorf("request from %s: status %d, want %d", from, got, status)
}
}
// saveUntilAnswered writes content to the state file at path, as an
// admin saves an edit of it, until a request to url from the client at
// from is answered with status. The file is written again before each
// request, since smallwebwaf may not watch its directory yet when it is
// first written. It waits as long as that takes, so that a slow test
// process cannot fail the test.
func saveUntilAnswered(t *testing.T, path, content, url, from string, status int) {
t.Helper()
for {
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
if statusFrom(t, url, from) == status {
return
}
time.Sleep(pollInterval)
}
}
// statusFrom returns the status a request to url from the client at
// from, as X-Forwarded-For names it, is answered with.
func statusFrom(t *testing.T, url, from string) int {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
http.NoBody) http.NoBody)
if err != nil { if err != nil {
@@ -474,7 +565,5 @@ func wantStatus(t *testing.T, url, from string, status int) {
_ = res.Body.Close() _ = res.Body.Close()
if res.StatusCode != status { return res.StatusCode
t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status)
}
} }
+256 -69
View File
@@ -1,14 +1,17 @@
// Package state keeps smallwebwaf's state in JSON files in // Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: // SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and // bans.json holds the bans, clients.json each client's counters and
// history, and lookups.json GeoJS's answers. Load reads them at start, and // history, and lookups.json GeoJS's answers. Load reads them at start,
// Run and WriteAll write them, each from a snapshot its part takes under // Watch takes in an admin's edit of one while smallwebwaf runs, and Run
// its own lock, so that no request waits on the disk. // and WriteAll write them. The disk is read and written outside the
// parts' locks, which are held only to take a snapshot or to put in what
// a file holds, so that no request waits on the disk.
package state package state
import ( import (
"bytes" "bytes"
"context" "context"
"crypto/sha256"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
@@ -17,10 +20,13 @@ import (
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"sync"
"time" "time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
) )
@@ -60,13 +66,26 @@ type Params struct {
// Now tells the time by which the counters' buckets run out, normally // Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC. // time.Now in UTC.
Now func() time.Time Now func() time.Time
// ProcessLog receives what was read, and the writes that fail. // ProcessLog receives what was read and taken in, the edits set aside,
// and the writes that fail.
ProcessLog *slog.Logger 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. // Files are the state files of a running smallwebwaf.
type Files struct { type Files struct {
params Params 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. // bansFile is bans.json, indented for an admin to read and edit.
@@ -117,41 +136,28 @@ func Load(params Params) (*Files, error) {
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err) return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
} }
var ( f := &Files{params: params, sums: map[string][sha256.Size]byte{}}
bansIn bansFile
clientsIn clientsFile
lookupsIn lookupsFile
)
err = errors.Join( bansRead, bansErr := f.read(bansJSON)
read(params.Dir, bansJSON, &bansIn), clientsRead, clientsErr := f.read(clientsJSON)
read(params.Dir, clientsJSON, &clientsIn), lookupsRead, lookupsErr := f.read(lookupsJSON)
read(params.Dir, lookupsJSON, &lookupsIn),
) err = errors.Join(bansErr, clientsErr, lookupsErr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
held := make([]bans.Ban, 0, len(bansIn.Bans))
for _, entry := range bansIn.Bans {
held = append(held, entry.ban())
}
params.Ledger.Load(held)
params.Limiter.Load(clientsIn.Clients, params.Now())
params.GeoJS.Load(lookupsIn.Lookups)
params.ProcessLog.Info("read the state files", "directory", params.Dir, params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", len(bansIn.Bans), "clients", len(clientsIn.Clients), "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
"lookups", len(lookupsIn.Lookups))
return &Files{params: params}, nil return f, nil
} }
// Run writes bans.json WriteDelay after a ban is made, with every ban // Run writes bans.json WriteDelay after a ban is made, with every ban
// made in between, and every file every CounterInterval, until ctx is // made in between, and every file every CounterInterval, until ctx is
// done. A write that fails is logged, and the file is written again at // done. A write that fails is logged, and the file is written again at
// its next write. // 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) { func (f *Files) Run(ctx context.Context) {
interval := time.NewTicker(f.params.CounterInterval) interval := time.NewTicker(f.params.CounterInterval)
defer interval.Stop() defer interval.Stop()
@@ -169,7 +175,7 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue: case <-bansDue:
bansDue = nil bansDue = nil
f.logFailure(f.writeBans()) f.logFailure(f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
f.logFailure(f.WriteAll()) f.logFailure(f.WriteAll())
} }
@@ -179,7 +185,50 @@ func (f *Files) Run(ctx context.Context) {
// WriteAll writes every state file, as smallwebwaf stops. A file that // WriteAll writes every state file, as smallwebwaf stops. A file that
// fails does not keep the others from being written. // fails does not keep the others from being written.
func (f *Files) WriteAll() error { func (f *Files) WriteAll() error {
return errors.Join(f.writeBans(), f.writeClients(), f.writeLookups()) return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON))
}
// 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:
f.takeInEdit(name)
}
case err = <-watcher.Errors:
f.params.ProcessLog.Warn("watching the state files failed",
"error", err.Error())
}
}
} }
// logFailure logs a write that failed. // logFailure logs a write that failed.
@@ -190,8 +239,170 @@ func (f *Files) logFailure(err error) {
} }
} }
// writeBans writes bans.json. // takeInEdit takes in an edit of the state file name and logs it, if the
func (f *Files) writeBans() error { // file has changed since smallwebwaf last read or wrote it and parses. A
// file that cannot be read or does not parse is left for its next write.
func (f *Files) takeInEdit(name string) {
f.mu.Lock()
defer f.mu.Unlock()
data, changed, err := f.readChanged(name)
if err != nil || !changed {
return
}
_, err = f.takeIn(name, data)
if err != nil {
return
}
// 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))
}
// 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)
}
// 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. 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) (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())
}
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)
}
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. An edit that does not parse is
// renamed to name.bad, for the admin to mend, and logged with where in
// the file the error is.
func (f *Files) writeFile(name string) error {
f.mu.Lock()
defer f.mu.Unlock()
path := filepath.Join(f.params.Dir, name)
data, changed, err := f.readChanged(name)
if err != nil {
return err
}
if changed {
_, err = f.takeIn(name, data)
if err == nil {
f.params.Metrics.StateFileEditTakenIn(name)
}
}
if err != nil {
renameErr := os.Rename(path, path+".bad")
if renameErr != nil {
return errors.Join(err, renameErr)
}
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse",
"file", path+".bad", "error", err.Error())
f.params.Metrics.StateFileEditSetAside(name)
}
data, err = f.encode(name)
if err != nil {
return fmt.Errorf("encode %s: %w", name, err)
}
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
}
// 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:
held := f.params.Ledger.Snapshot() held := f.params.Ledger.Snapshot()
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))} file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))}
@@ -201,30 +412,15 @@ func (f *Files) writeBans() error {
data, err := json.MarshalIndent(file, "", " ") data, err := json.MarshalIndent(file, "", " ")
if err != nil { if err != nil {
return fmt.Errorf("encode %s: %w", bansJSON, err) return nil, err
} }
return write(f.params.Dir, bansJSON, append(data, '\n')) return append(data, '\n'), nil
case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
default: // lookups.json
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
} }
// writeClients writes clients.json.
func (f *Files) writeClients() error {
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot())
if err != nil {
return fmt.Errorf("encode %s: %w", clientsJSON, err)
}
return write(f.params.Dir, clientsJSON, data)
}
// writeLookups writes lookups.json.
func (f *Files) writeLookups() error {
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
if err != nil {
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
}
return write(f.params.Dir, lookupsJSON, data)
} }
// newBanEntry returns ban as bans.json holds it. // newBanEntry returns ban as bans.json holds it.
@@ -377,28 +573,16 @@ func checkWritable(dir string) error {
return errors.Join(file.Close(), os.Remove(file.Name())) return errors.Join(file.Close(), os.Remove(file.Name()))
} }
// read reads the state file name in dir into file, a pointer to that // parse reads data, what the state file at path holds, into file, a
// file's struct, and checks its entries. A missing file leaves file as it // pointer to that file's struct, and checks its entries.
// is. func parse(path string, data []byte, file stateFile) error {
func read(dir, name string, file stateFile) error {
path := filepath.Join(dir, name)
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
if errors.Is(err, fs.ErrNotExist) {
return nil
}
if err != nil {
return err
}
// The version is read first, so that a file of another version is // 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. // refused for that, and not for an entry this version cannot read.
var header struct { var header struct {
Version int `json:"version"` Version int `json:"version"`
} }
err = json.Unmarshal(data, &header) err := json.Unmarshal(data, &header)
if err == nil && header.Version != version { if err == nil && header.Version != version {
err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d", err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d",
errVersion, header.Version, version) errVersion, header.Version, version)
@@ -451,7 +635,7 @@ func position(data []byte, err error) string {
// write writes data to the file name in dir so that a crash at any // 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 // 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 // temporary file in the same directory, which is synced and renamed over
// name, and then the directory is synced, so that the rename lasts. // name. syncDirectory must follow, so that the rename lasts.
func write(dir, name string, data []byte) error { func write(dir, name string, data []byte) error {
path := filepath.Join(dir, name) path := filepath.Join(dir, name)
temporary := path + ".tmp" temporary := path + ".tmp"
@@ -463,10 +647,13 @@ func write(dir, name string, data []byte) error {
if err != nil { if err != nil {
_ = os.Remove(temporary) _ = os.Remove(temporary)
}
return err 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 directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
if err != nil { if err != nil {
return err return err
+467 -12
View File
@@ -3,11 +3,16 @@ package state_test
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"io/fs"
"log/slog" "log/slog"
"net"
"net/http"
"net/http/httptest"
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"slices" "slices"
"strconv"
"strings" "strings"
"testing" "testing"
"testing/synctest" "testing/synctest"
@@ -15,6 +20,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/state" "sneak.berlin/go/smallwebwaf/internal/state"
) )
@@ -24,6 +30,13 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
// What the process log says once Watch watches the directory, and as
// it takes in an edit.
watching = "watching the state files for edits"
tookIn = "took in an edit of a state file"
// maxLogLines is how many lines of the process log wait for a test to
// read them.
maxLogLines = 64
) )
// permanentBansJSON is bans.json holding permanentBan. // permanentBansJSON is bans.json holding permanentBan.
@@ -286,7 +299,7 @@ func TestUnwritableDirectoryStopsTheStart(t *testing.T) {
} }
} }
// The two tests below run Run in a synctest bubble, where time is a clock // The three tests below run Run in a synctest bubble, where time is a clock
// of the test's own: time.Sleep moves it on at once, and synctest.Wait // of the test's own: time.Sleep moves it on at once, and synctest.Wait
// returns once Run waits for its next write, so that every write due by // returns once Run waits for its next write, so that every write due by
// then is on disk. // then is on disk.
@@ -298,7 +311,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
dir := t.TempDir() dir := t.TempDir()
params := newParams(dir) params := newParams(dir)
params.WriteDelay = 10 * time.Second params.WriteDelay = 10 * time.Second
run(t, load(t, params)) run(t, load(t, params).Run)
// A second ban, made while the first waits to be written, puts the // A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it. // write off no further, and is written with it.
@@ -343,7 +356,7 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
dir := t.TempDir() dir := t.TempDir()
params := newParams(dir) params := newParams(dir)
params.CounterInterval = time.Minute params.CounterInterval = time.Minute
run(t, load(t, params)) run(t, load(t, params).Run)
// The files are removed once written, so that each interval shows // The files are removed once written, so that each interval shows
// them written again. // them written again.
@@ -360,6 +373,44 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
}) })
} }
func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.WriteDelay = 10 * time.Second
run(t, load(t, params).Run)
// A ban, and bans.json written with it.
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
midnight(), bans.Notes{})
time.Sleep(params.WriteDelay)
synctest.Wait()
// A second ban is to be written WriteDelay later. Just before
// then, an admin saves bans.json with the first ban lifted and
// another added.
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
midnight(), bans.Notes{})
time.Sleep(params.WriteDelay - time.Nanosecond)
synctest.Wait()
edit(t, dir, bansJSON, permanentBansJSON)
// The write takes the edit in first, and writes it back. The second
// ban, made after the admin opened the file, is lost, as "Edits
// while running" in SPEC.md says.
time.Sleep(time.Nanosecond)
synctest.Wait()
if got := readFile(t, filepath.Join(dir, bansJSON)); got != permanentBansJSON {
t.Errorf("bans.json holds\n%s\nwant the edit", got)
}
wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()})
})
}
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) { func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel() t.Parallel()
@@ -402,26 +453,301 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
} }
} }
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { func TestWritesAreCountedInTheMetrics(t *testing.T) {
t.Parallel() t.Parallel()
dir := t.TempDir() dir := t.TempDir()
files := load(t, newParams(dir)) params := newParams(dir)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
// A directory named bans.json cannot be renamed over. err := files.WriteAll()
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) if err != nil {
t.Fatalf("write: %v", err)
}
const (
ofBans = `{file="bans.json"}`
ofClients = `{file="clients.json"}`
)
got := scrape(t, params)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 1)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofClients, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 0)
wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans,
float64(len(permanentBansJSON)))
written := metric(t, got, "smallwebwaf_state_file_last_write_timestamp_seconds"+ofBans)
if written < float64(time.Now().Add(-time.Hour).Unix()) {
t.Errorf("bans.json was last written at %v, not by that write", written)
}
// A directory in the way of bans.json's temporary file fails its next
// write, which leaves its size as it was, although it has a ban more.
err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700)
if err != nil { if err != nil {
t.Fatalf("mkdir: %v", err) t.Fatalf("mkdir: %v", err)
} }
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
err = files.WriteAll() err = files.WriteAll()
if err == nil { if err == nil {
t.Error("writing over a directory did not fail") t.Fatal("the write did not fail")
}
got = scrape(t, params)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 2)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofClients, 0)
wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans,
float64(len(permanentBansJSON)))
}
func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
files := load(t, newParams(dir))
// bans.json is a socket, which cannot be opened as a file, even by
// root, as the tests run in Docker, but which a rename could replace.
// Whether it holds an edit cannot be told, so it is left as it is.
socket, err := (&net.ListenConfig{}).Listen(t.Context(), "unix", path)
if err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
_ = socket.Close()
}()
err = files.WriteAll()
if err == nil {
t.Error("writing with bans.json unreadable did not fail")
}
info, err := os.Lstat(path)
if err != nil || info.Mode().Type() != fs.ModeSocket {
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
} }
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
} }
func TestEditOfEachFileTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// Each edit holds one entry, for a client the parts did not hold, and
// takes the place of everything the part held.
client := netip.MustParsePrefix("198.51.100.7/32")
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "198.51.100.7/32", `+
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON)
wantEqual(t, bansJSON, params.Ledger.Snapshot(),
[]bans.Ban{{Netblock: client, Start: midnight()}})
edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+
`{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`)
wantTakenIn(t, lines, dir, clientsJSON)
wantEqual(t, clientsJSON, params.Limiter.Snapshot(),
[]ratelimit.Client{{Client: client, History: ratelimit.History{Requests: 7}}})
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": [{"client": "198.51.100.7/32", `+
`"country": "FR", "answered": "2026-10-06T00:00:00Z"}]}`)
wantTakenIn(t, lines, dir, lookupsJSON)
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
}
func TestOwnWritesAreNotTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
fill(params)
files := load(t, params)
watch(t, files, lines)
// Every file is written while watched, and then lookups.json edited:
// the first edit taken in is that one.
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": []}`)
wantTakenIn(t, lines, dir, lookupsJSON)
}
func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
watch(t, load(t, params), lines)
client := netip.MustParseAddr("203.0.113.9")
// An entry added, as an admin writes it, bans its netblock.
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", `+
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned := params.Ledger.Check(client, midnight())
if !banned {
t.Error("the ban added to bans.json does not refuse")
}
// The entry removed lifts the ban.
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned = params.Ledger.Check(client, midnight())
if banned {
t.Error("the ban removed from bans.json still refuses")
}
}
func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Parallel()
// It ends a ban's entry with a comma.
const broken = "{\n \"version\": 1,\n \"bans\": [\n" +
" {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n"
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
lines := logInto(&params)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// While smallwebwaf runs, the broken edit is left as it is: an edit
// of clients.json, made after it and taken in, shows that it has been
// seen.
edit(t, dir, bansJSON, broken)
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, clientsJSON)
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
// The next write sets it aside, logged with where the error is, and
// writes bans.json again from what smallwebwaf still holds.
err = files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
line := lines.waitFor(t, "set aside an edit of a state file that does not parse")
message, _ := line["error"].(string)
if line["file"] != path+".bad" ||
!strings.HasPrefix(message, path+", line 4, column 39: ") {
t.Errorf("set aside with %v", line)
}
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
if got := readFile(t, path+".bad"); got != broken {
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
}
if got := readFile(t, path); got != permanentBansJSON {
t.Errorf("bans.json holds\n%s\nwant\n%s", got, permanentBansJSON)
}
}
func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
// One edit is taken in by the write of its file, before Watch runs,
// and one by Watch.
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
edit(t, dir, bansJSON, permanentBansJSON)
wantTakenIn(t, lines, dir, bansJSON)
wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2)
}
func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
files := load(t, params)
edit(t, dir, bansJSON, `{"version": 1, "bans": [`)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1)
}
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
err := os.Remove(dir)
if err != nil {
t.Fatalf("remove: %v", err)
}
// Watch returns at once.
files.Watch(t.Context())
line := lines.waitFor(t, "cannot watch the state files for edits")
if line["level"] != "ERROR" {
t.Errorf("logged as %v", line)
}
}
// midnight is the time of the tests' clock. // midnight is the time of the tests' clock.
func midnight() time.Time { func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC) return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
@@ -431,6 +757,7 @@ func midnight() time.Time {
// hold nothing yet. GeoJS is never asked. // hold nothing yet. GeoJS is never asked.
func newParams(dir string) state.Params { func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler) discard := slog.New(slog.DiscardHandler)
m := metrics.New(1)
return state.Params{ return state.Params{
Dir: dir, Dir: dir,
@@ -443,9 +770,12 @@ func newParams(dir string) state.Params {
MaxBans: 5000, MaxBans: 5000,
}), }),
Limiter: ratelimit.New(ratelimit.Limits{}), Limiter: ratelimit.New(ratelimit.Limits{}),
GeoJS: lookup.New(lookup.Params{Now: midnight, ProcessLog: discard}), GeoJS: lookup.New(lookup.Params{
Now: midnight, ProcessLog: discard, Metrics: m,
}),
Now: midnight, Now: midnight,
ProcessLog: discard, ProcessLog: discard,
Metrics: m,
} }
} }
@@ -512,15 +842,16 @@ func load(t *testing.T, params state.Params) *state.Files {
return files return files
} }
// run runs files' writes until the test ends. // run runs task, the Run or the Watch of state files, until the test
func run(t *testing.T, files *state.Files) { // ends.
func run(t *testing.T, task func(context.Context)) {
t.Helper() t.Helper()
ctx, stop := context.WithCancel(t.Context()) ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{}) stopped := make(chan struct{})
go func() { go func() {
files.Run(ctx) task(ctx)
close(stopped) close(stopped)
}() }()
@@ -530,6 +861,79 @@ func run(t *testing.T, files *state.Files) {
}) })
} }
// watch runs files' Watch until the test ends, and waits until it
// watches the directory.
func watch(t *testing.T, files *state.Files, lines processLog) {
t.Helper()
run(t, files.Watch)
lines.waitFor(t, watching)
}
// processLog receives the lines of a process log, each a JSON object, for
// a test to wait for.
type processLog chan string
// logInto has params' process log write its lines into a new processLog,
// and returns that.
func logInto(params *state.Params) processLog {
lines := make(processLog, maxLogLines)
params.ProcessLog = slog.New(slog.NewJSONHandler(lines, nil))
return lines
}
// 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 // never reached: nothing closes the log
}
// wantTakenIn waits for the next edit taken in, and checks that it is of
// the state file name in dir.
func wantTakenIn(t *testing.T, lines processLog, dir, name string) {
t.Helper()
line := lines.waitFor(t, tookIn)
if line["file"] != filepath.Join(dir, name) {
t.Fatalf("took in %v, want an edit of %s", line, name)
}
}
// edit writes content to the state file name in dir, as an admin saves an
// edit of it.
func edit(t *testing.T, dir, name, content string) {
t.Helper()
err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
}
// wantEqual checks that the entries read back from file are those // wantEqual checks that the entries read back from file are those
// written. // written.
func wantEqual[E comparable](t *testing.T, file string, got, want []E) { func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
@@ -633,3 +1037,54 @@ func wantEntries(t *testing.T, path, key string, want ...string) {
} }
} }
} }
// scrape returns the metrics of params, in the Prometheus text format.
func scrape(t *testing.T, params state.Params) string {
t.Helper()
recorder := httptest.NewRecorder()
params.Metrics.ServeHTTP(recorder,
httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody))
if recorder.Code != http.StatusOK {
t.Fatalf("the metrics were answered %d", recorder.Code)
}
return recorder.Body.String()
}
// metric returns the value of series in text, the metrics, such as
// smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test
// if there is no such series.
func metric(t *testing.T, text, series string) float64 {
t.Helper()
for line := range strings.Lines(text) {
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, text)
return 0
}
// wantMetric checks the value of series in text, the metrics, as metric
// reads it.
func wantMetric(t *testing.T, text, series string, want float64) {
t.Helper()
got := metric(t, text, series)
if got != want {
t.Errorf("%s is %v, want %v", series, got, want)
}
}
+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)
}
}