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

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

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

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

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

Model: opus-5-5
2026-10-06 14:18:13 +02:00
28 changed files with 2131 additions and 328 deletions
+148 -51
View File
@@ -13,23 +13,25 @@ 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 seven parts of https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are eight parts of
milestone 3: the static lists, the bans that broken rate limits lead to, the milestone 3: the static lists, the bans that broken rate limits lead to, the
JSON state files and the paths the rate limits do not count, which come next in JSON state files with your edits taken in while it runs and the paths the rate
the build order, `observe` mode, which comes a little later, and the metrics limits do not count, which come next in the build order, `observe` mode and the
rest of the request log's fields, which come a little later, and the metrics
endpoint and the header size and the idle time as settings, which come last in endpoint and the header size and the idle time as settings, which come last in
it. `smallwebwaf` passes each request to the app and the app's answer back, it. `smallwebwaf` passes each request to the app and the app's answer back,
unchanged, within its timeouts and size limits, works out each client's address, unchanged, within its timeouts and size limits, works out each client's address,
bans a client that sends too many requests, not counting those for the paths you bans a client that sends too many requests, not counting those for the paths you
choose, refuses a client that comes from a country you refuse or from a network choose, refuses a client that comes from a country you refuse or from a network
you refuse, lets the networks you choose through, keeps its bans, each client's you refuse, lets the networks you choose through, keeps its bans, each client's
counters and history, and GeoJS's answers in JSON files across restarts, writes counters and history, and GeoJS's answers in JSON files across restarts, takes
a JSON log line for every request, serves Prometheus metrics to a scraper that in your edits of those files while it runs, writes a JSON log line for every
holds the metrics token, and in `observe` mode passes on the requests it would request, serves Prometheus metrics to a scraper that holds the metrics token,
refuse, logging what it would have done with them. It comes as the image the and in `observe` mode passes on the requests it would refuse, logging what it
app's own image is built on. The rest of the design comes after that, in the would have done with them. It comes as the image the app's own image is built
order of the build order in [`SPEC.md`](SPEC.md). The survey of existing tools on. The rest of the design comes after that, in the order of the build order in
that led to the design is in [`EVALUATION.md`](EVALUATION.md). [`SPEC.md`](SPEC.md). The survey of existing tools that led to the design is in
[`EVALUATION.md`](EVALUATION.md).
## Getting started ## Getting started
@@ -66,7 +68,9 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
address outside `SWWAF_TRUSTED_PROXIES` is the client; if every address in it address outside `SWWAF_TRUSTED_PROXIES` is the client; if every address in it
is inside, the leftmost is, and with no header the peer is. The app sees what is inside, the leftmost is, and with no header the peer is. The app sees what
it would see from traefik directly: the same `Host`, the same it would see from traefik directly: the same `Host`, the same
`X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end. `X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end. It
also gets the request's id in `X-Request-ID`, the same id as in the request's
log line (see `request_id` in "Request log" below).
- Enforces the timeouts and the size limits below. A limit passed before the - Enforces the timeouts and the size limits below. A limit passed before the
response has started gets `smallwebwaf`'s own answer: `408` for a client too response has started gets `smallwebwaf`'s own answer: `408` for a client too
slow to send its request, `413` for a request body that is too large, `504` slow to send its request, `413` for a request body that is too large, `504`
@@ -81,13 +85,14 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
takes the client over one of the rate limits below is refused with takes the client over one of the rate limits below is refused with
`SWWAF_BAN_RESPONSE`, `403` by default, before anything reaches the app, and `SWWAF_BAN_RESPONSE`, `403` by default, before anything reaches the app, and
bans the client. A request whose path starts with one of bans the client. A request whose path starts with one of
`SWWAF_RATE_LIMIT_EXEMPT_PATHS` is neither counted nor refused by the rate `SWWAF_RATE_LIMIT_EXEMPT_PATHS`, as that setting below describes, is neither
limits; the static lists, bans and the country lists still apply to it. A counted nor refused by the rate limits; the static lists, bans and the country
client is one IPv4 address, or one IPv6 /64, since one abuser usually holds a lists still apply to it. A client is one IPv4 address, or one IPv6 /64, since
whole /64. Each window is counted in two fixed buckets, the earlier one one abuser usually holds a whole /64. Each window is counted in two fixed
weighted by how much of it the window still covers. At most 20,000 clients are buckets, the earlier one weighted by how much of it the window still covers.
kept, the least recently seen dropped first, with their history, and a restart At most 20,000 clients are kept, the least recently seen dropped first, with
gives no client a fresh allowance (see "State files" below). their history, and a restart gives no client a fresh allowance (see "State
files" below).
- Bans a client that breaks a rate limit, as "Bans" in [`SPEC.md`](SPEC.md) - Bans a client that breaks a rate limit, as "Bans" in [`SPEC.md`](SPEC.md)
describes: the first ban lasts an hour, and a limit broken again within a day describes: the first ban lasts an hour, and a limit broken again within a day
of a ban ending bans for three times as long as that ban, so 1, 3, 9, 27 and of a ban ending bans for three times as long as that ban, so 1, 3, 9, 27 and
@@ -103,9 +108,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
@@ -156,6 +161,11 @@ it, and the effective settings are logged at start.
- `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens. - `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens.
- `SWWAF_UPSTREAM_URL` (default `http://127.0.0.1:8081`): the app, as `http` or - `SWWAF_UPSTREAM_URL` (default `http://127.0.0.1:8081`): the app, as `http` or
`https`, a host and an optional port, and nothing more. `https`, a host and an optional port, and nothing more.
- `SWWAF_INSTANCE_NAME` (default: the host's name, which docker sets to the
first 12 characters of the container's id unless the deployment names one):
the name each request log line gives as `instance`. Set it, for example to
`fsn1app1/gitea`, for a name that stays the same when a deploy replaces the
container, and that tells instances apart when several log to one place.
- `SWWAF_MODE` (default `enforce`): `enforce`, or `observe` to pass on the - `SWWAF_MODE` (default `enforce`): `enforce`, or `observe` to pass on the
requests `smallwebwaf` would refuse and log what it would have done (see "What requests `smallwebwaf` would refuse and log what it would have done (see "What
it does so far" above). it does so far" above).
@@ -196,15 +206,17 @@ it, and the effective settings are logged at start.
page makes a few hundred requests and several people often share one address. page makes a few hundred requests and several people often share one address.
- `SWWAF_RATE_LIMIT_EXEMPT_PATHS` (default empty): path prefixes whose requests - `SWWAF_RATE_LIMIT_EXEMPT_PATHS` (default empty): path prefixes whose requests
the rate limits neither count nor refuse, such as `/assets/` for static the rate limits neither count nor refuse, such as `/assets/` for static
assets; each starts with `/`. A prefix is compared, character for character, assets; each starts with `/`. A request whose path, percent-decoded, contains
with the start of the path the app will act on: the request's path, before any `..` anywhere or a backslash, or whose path as sent holds an encoded slash
query string, percent-decoded, with its `.` and `..` segments and repeated (`%2F` or `%2f`), is never exempt, since the app may act on it as a path
slashes resolved and without a trailing slash, which is not always what the outside every prefix: `/assets/..%2Flogin` as `/login`. Any other request is
request log's `path` shows. `/assets/` matches `/assets/app.js`, exempt when its path as sent, the path the app receives, before any query
`/assets//img/logo.png` and `/static/../assets/app.js`, but not `/assets/` string and not percent-decoded, starts with a prefix, character for character.
itself, `/assets`, `/Assets/app.js`, `/static/assets/app.js` or `/assets/` matches `/assets/app.js` and `/assets/`, but not `/assets`,
`/assets/..%2Flogin`, which is `/login`. A prefix is written without `/Assets/app.js`, `/%61ssets/app.js`, `/static/assets/app.js`,
percent-encoding, and there are no wildcards: `*` is a character like any `/static/../assets/app.js` or `/assets%2Fapp.js`. A character the client sends
percent-encoded, such as a space, is written percent-encoded in a prefix, as
in `/my%20files/`, and there are no wildcards: `*` is a character like any
other. other.
- `SWWAF_DENIED_COUNTRIES` (default empty): countries whose clients are refused, - `SWWAF_DENIED_COUNTRIES` (default empty): countries whose clients are refused,
for example `cn,ru,kp`. for example `cn,ru,kp`.
@@ -235,6 +247,13 @@ 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_LOG_REQUEST_HEADERS` (default
`accept,accept-language,accept-encoding,content-type,origin,range`): the
request headers whose values the request log gives, in either case.
`Authorization`, `Cookie` and `Set-Cookie` are never logged, even when listed
(see "Request log" below). An entry naming `Host` or `Transfer-Encoding` stops
the start, since Go's HTTP server takes both out of the request; the request's
host is the field `host`.
- `SWWAF_METRICS_TOKEN` (default unset): the token a scraper sends for the - `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 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 shorter than 32 characters stops the start. The settings logged at start show
@@ -264,18 +283,42 @@ GeoJS are kept, for 7 days each.
refused ones included: refused ones included:
``` ```
{"type":"request","time":"2026-10-03T12:00:00.123Z","client_ip":"203.0.113.9","peer_ip":"172.18.0.2","country":"DE","method":"GET","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"upstream_status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","action":"forward","duration_total":3.217,"duration_upstream_total":3.104} {"type":"request","time":"2026-10-03T12:00:00.123Z","instance":"fsn1app1/gitea","client_ip":"203.0.113.9","method":"GET","scheme":"https","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","request_id":"7Q2NHZ4KJ3VXW5YB6R3MEFTD2A","peer_ip":"172.18.0.2","forwarded_for":"203.0.113.9","client_group":"203.0.113.9/32","country":"DE","request_headers":{"accept":"*/*"},"response_content_type":"text/html; charset=utf-8","upstream_status":200,"action":"forward","counts":{"minute":1,"hour":12,"day":40},"duration_total":3.217,"duration_checks":0.041,"duration_upstream_connect":0.052,"duration_upstream_first_byte":2.874,"duration_upstream_total":3.104}
``` ```
- `time` is when the request arrived, in UTC. `peer_ip` is the TCP peer, A field that does not apply to a request is left out of its line, apart from
normally traefik. `path` and `query` are as the client sent them. `type`, the fields from `time` to `user_agent`, `request_id`, `peer_ip`,
`client_group`, `country`, `action` and `duration_total`, which every line has.
- `time` is when the request arrived, in UTC. `instance` is
`SWWAF_INSTANCE_NAME`. `scheme` is the `X-Forwarded-Proto` a trusted proxy
sent, and otherwise `http`. `path` and `query` are as the client sent them.
- `request_id` is the `X-Request-ID` a trusted proxy sent, or a new random one
of 26 letters and digits when it sent none, or when the peer is not a trusted
proxy. A request passed to the app takes it there in `X-Request-ID`.
- `peer_ip` is the TCP peer, normally traefik. `forwarded_for` is the
`X-Forwarded-For` header as received, several lines of it joined with `, `.
`client_group` is the client as the rate limits count it: its IPv4 address as
a /32, or the /64 of its IPv6 address.
- `country` is the client's country as GeoJS places it. It is empty with neither - `country` is the client's country as GeoJS places it. It is empty with neither
country list set, for a client in `SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS`, for country list set, for a client in `SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS`, for
a client on a private, loopback or link-local address, when GeoJS cannot place a client on a private, loopback or link-local address, when GeoJS cannot place
the client or has not answered in time, and for a request whose client a ban the client or has not answered in time, and for a request whose client a ban
covers, even when the client's country is known. covers, even when the client's country is known.
- `content_type` is the request's `Content-Type`, and `content_length` the
length the request announced for its body, which is left out for none or zero.
- `request_headers` are the request's headers that `SWWAF_LOG_REQUEST_HEADERS`
names, by name in lower case, several lines of one joined with `, `.
`Authorization`, `Cookie` and `Set-Cookie` are never among them, whatever the
setting says: `has_authorization` and `has_cookie` are there instead, and
true, when the request has an `Authorization` or a `Cookie` header.
- `websocket` is there, and true, when the app switched the connection to
another protocol, as it does for a WebSocket.
- `status` is what the client was sent, `0` if nothing was; `upstream_status` is - `status` is what the client was sent, `0` if nothing was; `upstream_status` is
what the app answered, and is left out when the app did not answer. what the app answered, and is left out when the app did not answer.
- `response_content_type`, `cache_control` and `location` are the
`Content-Type`, `Cache-Control` and `Location` headers of the answer: the
app's, as passed on, or those of `smallwebwaf`'s own answer.
- `request_bytes` and `response_bytes` count body bytes. - `request_bytes` and `response_bytes` count body bytes.
- `action` is `forward` for a request passed to the app, `denied` for one - `action` is `forward` for a request passed to the app, `denied` for one
refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused
@@ -291,6 +334,16 @@ refused ones included:
`banned`, `country_denied` or `rate_limited`. `action` then names what was `banned`, `country_denied` or `rate_limited`. `action` then names what was
done: `forward` for a request passed to the app, and another action, such as done: `forward` for a request passed to the app, and another action, such as
`too_large`, for one a size or time limit refused. `too_large`, for one a size or time limit refused.
- `counts` gives the client's requests in the minute, the hour and the day as
the rate limits count them, this request included: in each window, those in
the bucket under way and a share of those in the bucket before, so a count can
have a fraction. For a request that broke a limit, they are the counts that
broke it. It is left out for a request the rate limits do not count: the
health check, one from a client in `SWWAF_ALLOW_NETS` or
`SWWAF_RATE_LIMIT_EXEMPT_NETS`, one for a path that
`SWWAF_RATE_LIMIT_EXEMPT_PATHS` exempts, and one that `SWWAF_DENY_NETS`, a ban
or the country lists refuse, or would refuse in `observe` mode. The byte
totals come with the byte limits.
- `limit_hit` is there for a request that broke a rate limit, and names the - `limit_hit` is there for a request that broke a rate limit, and names the
window whose limit it went over: `minute`, `hour` or `day`, the shortest if it window whose limit it went over: `minute`, `hour` or `day`, the shortest if it
went over several. `offence` is then `limit`. went over several. `offence` is then `limit`.
@@ -298,10 +351,19 @@ refused ones included:
or in `observe` mode would have been refused under one, and gives when the ban or in `observe` mode would have been refused under one, and gives when the ban
ends, in the same form as `time`, or `permanent`. ends, in the same form as `time`, or `permanent`.
- `aborted` is there, and true, when the client went away early. - `aborted` is there, and true, when the client went away early.
- `duration_total` and `duration_upstream_total` are in milliseconds. - The timings are in milliseconds, to the microsecond. `duration_total` runs
from when the request's headers had been read to when its line is written, and
`duration_checks` over the same start to when the checks were done; the health
check runs none, and its line has no `duration_checks`.
`duration_upstream_connect`, `duration_upstream_first_byte` and
`duration_upstream_total` are there for a request passed to the app, and run
from when it was handed to the app: until there was a connection to it, new or
kept open from an earlier request, until the first byte of its answer arrived,
and until the end. The first two are left out when that never happened, as for
an app that cannot be reached.
No body and no other header is logged. `smallwebwaf`'s own messages (start, the No body is logged, and no header but those above. `smallwebwaf`'s own messages
settings, stop, errors) share the stream as JSON lines marked (start, the settings, stop, errors) share the stream as JSON lines marked
`"type":"process"`. `"type":"process"`.
Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and
@@ -350,9 +412,42 @@ 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 ## Metrics
@@ -391,7 +486,10 @@ other request. No metric carries a client's address.
- `smallwebwaf_state_file_writes_total`, - `smallwebwaf_state_file_writes_total`,
`smallwebwaf_state_file_write_failures_total`, `smallwebwaf_state_file_write_failures_total`,
`smallwebwaf_state_file_last_write_timestamp_seconds` and `smallwebwaf_state_file_last_write_timestamp_seconds` and
`smallwebwaf_state_file_size_bytes`, by `file`. `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. - Go's own `go_` metrics and the process's `process_` metrics.
The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request
@@ -499,9 +597,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
@@ -693,8 +791,8 @@ addresses are never sent to GeoJS.
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
@@ -706,8 +804,9 @@ 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, and netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and
`github.com/prometheus/client_golang` keeps the metrics and serves them. The `github.com/prometheus/client_golang` keeps the metrics and serves them, and
country codes are the list in `internal/config/config.go`. `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
@@ -750,10 +849,8 @@ so that they run in minimal containers.
## TODO ## TODO
- The rest of milestone 3: taking in an admin's edits to the state files - The rest of the design, in the order of the build order in
(https://git.eeqj.de/sneak/smallwebwaf/issues/68), exemptions and the rest of [`SPEC.md`](SPEC.md).
the request log's fields; then the rest of the design, in the order of the
build order in [`SPEC.md`](SPEC.md).
## Documents ## Documents
+1
View File
@@ -3,6 +3,7 @@ module sneak.berlin/go/smallwebwaf
go 1.26.0 go 1.26.0
require ( require (
github.com/fsnotify/fsnotify v1.10.1
github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/hashicorp/golang-lru/v2 v2.0.7
github.com/prometheus/client_golang v1.24.1 github.com/prometheus/client_golang v1.24.1
) )
+2
View File
@@ -4,6 +4,8 @@ github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UF
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= 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 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
+50 -24
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.
@@ -176,9 +176,23 @@ func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
return *ban, true return *ban, true
} }
// 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
@@ -192,12 +206,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()
@@ -282,21 +306,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()
@@ -318,11 +346,9 @@ func (l *Ledger) active(client netip.Addr, now time.Time) *Ban {
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] return ban
if last.ActiveAt(now) {
return last
} }
} }
@@ -358,8 +384,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
+104
View File
@@ -130,6 +130,75 @@ 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)
client := netip.MustParseAddr("203.0.113.9")
ban, banned := ledger.Find(client, now)
if !banned || !ban.Permanent() {
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
}
ban, banned = ledger.Check(client, now)
if !banned || !ban.Permanent() {
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
banned, ban)
}
// A limit broken now makes no shorter ban over the permanent one.
ban = ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
}
}
func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
t.Parallel()
// A 9-hour ban smallwebwaf made, the third in a row, and an admin's
// 1-hour ban added to bans.json over it, with no 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 +220,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()
+66 -2
View File
@@ -12,6 +12,7 @@ import (
"net/http" "net/http"
"net/netip" "net/netip"
"net/url" "net/url"
"os"
"path/filepath" "path/filepath"
"slices" "slices"
"strconv" "strconv"
@@ -27,6 +28,10 @@ type Config struct {
ListenAddr string ListenAddr string
// UpstreamURL is the app (SWWAF_UPSTREAM_URL). // UpstreamURL is the app (SWWAF_UPSTREAM_URL).
UpstreamURL *url.URL UpstreamURL *url.URL
// InstanceName is the name each request log line gives as instance
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets
// to the first 12 characters of the container's id.
InstanceName string
// Observe is true in observe mode, when SWWAF_MODE is observe rather // Observe is true in observe mode, when SWWAF_MODE is observe rather
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country // than enforce: a request that SWWAF_DENY_NETS, a ban, the country
// lists or a rate limit would refuse is passed to the app instead, and // lists or a rate limit would refuse is passed to the app instead, and
@@ -113,6 +118,9 @@ type Config struct {
StateDir string StateDir string
StateWriteDelay time.Duration StateWriteDelay time.Duration
StateCounterInterval time.Duration StateCounterInterval time.Duration
// LogRequestHeaders are the request headers whose values the request
// log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS).
LogRequestHeaders []string
// MetricsToken is the bearer token a scraper sends for the metrics // MetricsToken is the bearer token a scraper sends for the metrics
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off. // (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
// MetricsTopN is how many countries get series of their own in the // MetricsTopN is how many countries get series of their own in the
@@ -159,6 +167,11 @@ var (
"such as http://127.0.0.1:8081") "such as http://127.0.0.1:8081")
errNotCountry = errors.New( errNotCountry = errors.New(
"is not a two-letter country code such as de or kp") "is not a two-letter country code such as de or kp")
errNotHeaderName = errors.New(
"is not a header name such as accept-language")
errHeaderTakenOut = errors.New(
"is taken out of every request by Go's HTTP server, so it can never " +
"be logged")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too") errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K") errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New( errNotDurationAboveZero = errors.New(
@@ -181,9 +194,11 @@ var (
// that is set but invalid is an error that names it. // that is set but invalid is an error that names it.
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) { func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv} env := &environment{lookupEnv: lookupEnv}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg := &Config{ cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"), ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"), UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
Observe: env.observe("SWWAF_MODE", "enforce"), Observe: env.observe("SWWAF_MODE", "enforce"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges), TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"), ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
@@ -214,8 +229,10 @@ 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"), LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"), "accept,accept-language,accept-encoding,content-type,origin,range"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
} }
for _, country := range cfg.ExclusivelyAllowedCountries { for _, country := range cfg.ExclusivelyAllowedCountries {
@@ -356,6 +373,15 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries return countries
} }
// headerNames reads a setting that is a list of header names, and
// returns them in lower case.
func (e *environment) headerNames(name, defaultValue string) []string {
headers, err := parseHeaderNames(e.value(name, defaultValue))
e.check(name, err)
return headers
}
// durationNotOff reads a setting that is a duration and, unlike a // durationNotOff reads a setting that is a duration and, unlike a
// timeout, cannot be off. // timeout, cannot be off.
func (e *environment) durationNotOff(name, defaultValue string) time.Duration { func (e *environment) durationNotOff(name, defaultValue string) time.Duration {
@@ -699,6 +725,44 @@ func parseCountries(value string) ([]string, error) {
return countries, nil return countries, nil
} }
// headerNameChars are the characters RFC 9110 allows in a header name:
// letters, digits and these marks.
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
"0123456789!#$%&'*+-.^_`|~"
// parseHeaderNames reads a comma-separated list of header names in either
// case, and returns them in lower case. Host and Transfer-Encoding are
// refused: Go's HTTP server takes them out of the request's headers.
func parseHeaderNames(value string) ([]string, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
headers := make([]string, 0, len(items))
for _, item := range items {
for _, char := range item {
if !strings.ContainsRune(headerNameChars, char) {
return nil, fmt.Errorf("%q %w", item, errNotHeaderName)
}
}
header := strings.ToLower(item)
switch header {
case "host":
return nil, fmt.Errorf("%q %w; the request's host is the field host",
item, errHeaderTakenOut)
case "transfer-encoding":
return nil, fmt.Errorf("%q %w", item, errHeaderTakenOut)
}
headers = append(headers, header)
}
return headers, nil
}
// parseListenAddr checks an address to listen on: an optional host and a // parseListenAddr checks an address to listen on: an optional host and a
// port number. // port number.
func parseListenAddr(value string) (string, error) { func parseListenAddr(value string) (string, error) {
+66 -7
View File
@@ -6,6 +6,7 @@ import (
"log/slog" "log/slog"
"maps" "maps"
"net/netip" "net/netip"
"os"
"slices" "slices"
"strings" "strings"
"testing" "testing"
@@ -48,8 +49,14 @@ const (
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N" metricsTopN = "SWWAF_METRICS_TOP_N"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
) )
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," +
"content-type,origin,range"
// token is a token of 32 characters, the shortest allowed. // token is a token of 32 characters, the shortest allowed.
const token = "0123456789abcdef0123456789abcdef" const token = "0123456789abcdef0123456789abcdef"
@@ -122,6 +129,18 @@ func TestDefaults(t *testing.T) {
wantCountries(t, deniedCountries, cfg.DeniedCountries) wantCountries(t, deniedCountries, cfg.DeniedCountries)
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries) wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries)
hostname, err := os.Hostname()
if err != nil || hostname == "" || cfg.InstanceName != hostname {
t.Errorf("%s is %q, want the host's name %q (%v)", instanceName,
cfg.InstanceName, hostname, err)
}
wantHeaders := strings.Split(defaultLogRequestHeaders, ",")
if !slices.Equal(cfg.LogRequestHeaders, wantHeaders) {
t.Errorf("%s gave %v, want %v", logRequestHeaders, cfg.LogRequestHeaders,
wantHeaders)
}
if len(cfg.RateLimitExemptPaths) != 0 { if len(cfg.RateLimitExemptPaths) != 0 {
t.Errorf("%s gave %v, want none", rateLimitExemptPaths, cfg.RateLimitExemptPaths) t.Errorf("%s gave %v, want none", rateLimitExemptPaths, cfg.RateLimitExemptPaths)
} }
@@ -222,6 +241,21 @@ func TestPathPrefixNotStartingWithSlashStopsTheStart(t *testing.T) {
} }
} }
func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
instanceName: "fsn1app1/gitea",
logRequestHeaders: " Accept , X-Custom",
})
if cfg.InstanceName != "fsn1app1/gitea" ||
!slices.Equal(cfg.LogRequestHeaders, []string{"accept", "x-custom"}) {
t.Errorf("%s is %q and %s gives %v", instanceName, cfg.InstanceName,
logRequestHeaders, cfg.LogRequestHeaders)
}
}
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) { func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -323,9 +357,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
for _, tc := range []struct{ name, value string }{ for _, tc := range []struct{ name, value string }{
{listenAddr, "8080"}, {listenAddr, "8080"}, {listenAddr, ":http"}, {listenAddr, ":65536"},
{listenAddr, ":http"},
{listenAddr, ":65536"},
{upstreamURL, "127.0.0.1:8081"}, {upstreamURL, "127.0.0.1:8081"},
{upstreamURL, "ftp://127.0.0.1:8081"}, {upstreamURL, "ftp://127.0.0.1:8081"},
{upstreamURL, "http://"}, {upstreamURL, "http://"},
@@ -343,8 +375,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowNets, "192.0.2.0/24,monitoring"}, {allowNets, "192.0.2.0/24,monitoring"},
{rateLimitExemptNets, "2001:db8::/129"}, {rateLimitExemptNets, "2001:db8::/129"},
{denyNets, "198.51.100.0/24,"}, {denyNets, "198.51.100.0/24,"},
{clientRequestTimeout, "60"}, {clientRequestTimeout, "60"}, {clientRequestTimeout, ""},
{clientRequestTimeout, ""},
{clientIdleTimeout, "0s"}, {clientIdleTimeout, "0s"},
{clientIdleTimeout, "2 minutes"}, {clientIdleTimeout, "2 minutes"},
{clientResponseTimeout, "1y"}, {clientResponseTimeout, "1y"},
@@ -361,8 +392,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{rateLimitPerMinute, "1K"}, {rateLimitPerMinute, "1K"},
{rateLimitPerHour, "0"}, {rateLimitPerHour, "0"},
{rateLimitPerHour, "1.5"}, {rateLimitPerHour, "1.5"},
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
{rateLimitPerDay, "lots"},
{rateLimitExemptPaths, "/assets/,,/static/"}, {rateLimitExemptPaths, "/assets/,,/static/"},
{deniedCountries, "nk"}, {deniedCountries, "nk"},
{deniedCountries, "kp,,ir"}, {deniedCountries, "kp,,ir"},
@@ -386,6 +416,10 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{stateWriteDelay, off}, {stateWriteDelay, "0s"}, {stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"}, {stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"}, {metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
{logRequestHeaders, "host"}, {logRequestHeaders, "accept,Host"},
{logRequestHeaders, "transfer-encoding"}, {logRequestHeaders, "TRANSFER-ENCODING"},
} { } {
t.Run(tc.name+"="+tc.value, func(t *testing.T) { t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -402,6 +436,27 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
} }
} }
func TestHostOrTransferEncodingStopsTheStart(t *testing.T) {
t.Parallel()
// Only Host's message points to the field host.
for value, want := range map[string]string{
"Host": `"Host" is taken out of every request by Go's HTTP server, ` +
"so it can never be logged; the request's host is the field host",
"transfer-encoding": `"transfer-encoding" is taken out of every ` +
"request by Go's HTTP server, so it can never be logged",
} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{logRequestHeaders: value}.lookupEnv)
if err == nil || err.Error() != logRequestHeaders+": "+want {
t.Errorf("error %v, want %s: %s", err, logRequestHeaders, want)
}
})
}
}
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) { func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel() t.Parallel()
@@ -453,6 +508,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Fatalf("decode %s: %v", out.Bytes(), err) t.Fatalf("decode %s: %v", out.Bytes(), err)
} }
hostname, _ := os.Hostname()
want := map[string]string{ want := map[string]string{
listenAddr: ":8080", listenAddr: ":8080",
upstreamURL: "http://127.0.0.1:8081", upstreamURL: "http://127.0.0.1:8081",
@@ -486,6 +543,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
stateCounterInterval: "15m", stateCounterInterval: "15m",
metricsToken: "", metricsToken: "",
metricsTopN: "50", metricsTopN: "50",
instanceName: hostname,
logRequestHeaders: defaultLogRequestHeaders,
} }
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)
+7 -5
View File
@@ -197,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 {
+20
View File
@@ -44,6 +44,8 @@ type Metrics struct {
stateFileWriteFailures *prometheus.CounterVec stateFileWriteFailures *prometheus.CounterVec
stateFileLastWrite *prometheus.GaugeVec stateFileLastWrite *prometheus.GaugeVec
stateFileSize *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. // New returns the metrics, with the Go runtime's and the process's own.
@@ -105,6 +107,11 @@ func New(topN int) *Metrics {
"When each state file was last written, in seconds since 1970.", byFile), "When each state file was last written, in seconds since 1970.", byFile),
stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes", stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes",
"The size of each state file, as it was last written.", byFile), "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.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
@@ -120,6 +127,7 @@ func New(topN int) *Metrics {
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered, m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures, m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize, m.stateFileLastWrite, m.stateFileSize,
m.stateFileEditsTakenIn, m.stateFileEditsSetAside,
) )
return m return m
@@ -230,6 +238,18 @@ func (m *Metrics) StateFileWritten(name string, size int, err error) {
m.stateFileSize.WithLabelValues(name).Set(float64(size)) 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 // statusClass returns the class of status, such as 2xx, or none when no
// status was sent. // status was sent.
func statusClass(status int) string { func statusClass(status int) string {
+8 -5
View File
@@ -30,14 +30,17 @@ func (rq *request) banned(now time.Time) bool {
return banned return banned
} }
// limitBroken counts the request for the rate limits at now, and reports // limitBroken counts the request for the rate limits at now, notes the
// whether it takes the client over one. In enforce mode such a request // client's counts for the log line, and reports whether the request takes
// bans the client's netblock, and sets the client's counters back to // the client over a limit. In enforce mode such a request bans the
// zero; in observe mode it does neither. // client's netblock, and sets the client's counters back to zero; in
// observe mode it does neither.
func (rq *request) limitBroken(now time.Time) bool { func (rq *request) limitBroken(now time.Time) bool {
group := clientGroup(rq.client) group := clientGroup(rq.client)
hit, over := rq.h.limiter.Count(group, now) counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts
if !over { if !over {
return false return false
} }
+28
View File
@@ -1,6 +1,7 @@
package proxy package proxy
import ( import (
"crypto/rand"
"net/http" "net/http"
"net/netip" "net/netip"
"slices" "slices"
@@ -48,6 +49,33 @@ func clientAddress(
return client return client
} }
// requestIDHeader carries the request's id, from traefik and to the app.
const requestIDHeader = "X-Request-ID"
// requestID is the request's id: the one a trusted proxy sent, or a new
// random one. A peer outside the trusted proxies did not come through
// traefik, so the id it sends is its own claim, and is replaced.
func requestID(r *http.Request, peerTrusted bool) string {
id := r.Header.Get(requestIDHeader)
if !peerTrusted || id == "" {
id = rand.Text()
}
return id
}
// scheme is how the client reached traefik, as a trusted proxy says in
// X-Forwarded-Proto, or otherwise http, the only scheme smallwebwaf
// serves.
func scheme(r *http.Request, peerTrusted bool) string {
proto := r.Header.Get("X-Forwarded-Proto")
if !peerTrusted || proto == "" {
return "http"
}
return proto
}
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client. // ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
const ipv6GroupPrefix = 64 const ipv6GroupPrefix = 64
+17 -13
View File
@@ -14,10 +14,14 @@ const (
appHost = "app.example" appHost = "app.example"
// client is the client's address, as a proxy names it. // client is the client's address, as a proxy names it.
client = "203.0.113.9" client = "203.0.113.9"
// forwardedFor is the header that lists the client and its proxies. // forwardedFor is the header that lists the client and its proxies,
forwardedFor = "X-Forwarded-For" // and forwardedProto the one that gives the scheme the client used.
// secure is the scheme a client reached traefik with. forwardedFor = "X-Forwarded-For"
forwardedProto = "X-Forwarded-Proto"
// secure is the scheme a client reached traefik with, and plain the
// one smallwebwaf serves.
secure = "https" secure = "https"
plain = "http"
) )
// appHeaders is what the app tells about the headers it received. // appHeaders is what the app tells about the headers it received.
@@ -65,13 +69,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
func clientAddressCases() []clientAddressCase { func clientAddressCases() []clientAddressCase {
trusted := map[string]string{trustedProxies: trustLocalhost} trusted := map[string]string{trustedProxies: trustLocalhost}
forged := http.Header{ forged := http.Header{
forwardedFor: {client}, forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"}, "X-Forwarded-Host": {"forged.example"},
"X-Forwarded-Proto": {secure}, forwardedProto: {secure},
"X-Real-Ip": {client}, "X-Real-Ip": {client},
} }
replaced := appHeaders{ replaced := appHeaders{
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http", ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
} }
return []clientAddressCase{{ return []clientAddressCase{{
@@ -87,10 +91,10 @@ func clientAddressCases() []clientAddressCase {
"outside the trusted proxies from the right", "outside the trusted proxies from the right",
env: trusted, env: trusted,
header: http.Header{ header: http.Header{
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"}, forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost}, "X-Forwarded-Host": {appHost},
"X-Forwarded-Proto": {secure}, forwardedProto: {secure},
"X-Real-Ip": {client}, "X-Real-Ip": {client},
}, },
wantClient: client, wantClient: client,
wantApp: appHeaders{ wantApp: appHeaders{
@@ -138,7 +142,7 @@ func requestWithHeaders(
Host: r.Host, Host: r.Host,
ForwardedFor: r.Header.Get(forwardedFor), ForwardedFor: r.Header.Get(forwardedFor),
ForwardedHost: r.Header.Get("X-Forwarded-Host"), ForwardedHost: r.Header.Get("X-Forwarded-Host"),
ForwardedProto: r.Header.Get("X-Forwarded-Proto"), ForwardedProto: r.Header.Get(forwardedProto),
RealIP: r.Header.Get("X-Real-IP"), RealIP: r.Header.Get("X-Real-IP"),
}) })
}) })
+12 -3
View File
@@ -21,14 +21,18 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
// the last one would have it refused. // the last one would have it refused.
addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"}) addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"})
const healthChecks = 3 const (
healthChecks = 3
contentType = "text/plain; charset=utf-8"
)
for range healthChecks { for range healthChecks {
got := get(t, addr, proxy.HealthPath) got := get(t, addr, proxy.HealthPath)
wantStatus(t, got, http.StatusOK) wantStatus(t, got, http.StatusOK)
if string(got.body) != "ok\n" { if string(got.body) != "ok\n" || got.header.Get("Content-Type") != contentType {
t.Errorf("health endpoint answered %q, want ok", got.body) t.Errorf("health endpoint answered %q with Content-Type %q, want ok "+
"with %q", got.body, got.header.Get("Content-Type"), contentType)
} }
} }
@@ -37,6 +41,11 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
lines := out.requestLines(t, healthChecks+1) lines := out.requestLines(t, healthChecks+1)
for _, line := range lines[:healthChecks] { for _, line := range lines[:healthChecks] {
wantLine(t, line, http.StatusOK, requestlog.ActionAdmin) wantLine(t, line, http.StatusOK, requestlog.ActionAdmin)
if line.ResponseContentType != contentType {
t.Errorf("health check's log line has response_content_type %q, "+
"want %q", line.ResponseContentType, contentType)
}
} }
wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward) wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward)
+30 -15
View File
@@ -6,6 +6,8 @@ import (
"errors" "errors"
"io" "io"
"net/http" "net/http"
"os"
"reflect"
"slices" "slices"
"strings" "strings"
"sync/atomic" "sync/atomic"
@@ -14,6 +16,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -115,27 +118,34 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
} }
} }
// wantRequestFields checks the log line's fields about the request. // wantRequestFields checks the log line's fields about the request. Its
// time, its id and its timings are checked only for being there.
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper() t.Helper()
want := requestlog.Line{ hostname, _ := os.Hostname()
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery, want := withTimings(line, requestlog.Line{
Protocol: "HTTP/1.1", Status: http.StatusTeapot, Type: requestType, Time: line.Time, Instance: hostname,
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent), ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent),
ResponseBytes: int64(received), UserAgent: "test-agent", ResponseBytes: int64(received), UserAgent: "test-agent",
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal, RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
DurationUpstreamTotal: line.DurationUpstreamTotal, ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
} UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
if line.Line != want { Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
} }
_, err := time.Parse(time.RFC3339, line.Time) _, err := time.Parse(time.RFC3339, line.Time)
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 { if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
t.Errorf("log line has time %q and durations %v and %v", line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
line.Time, line.DurationTotal, line.DurationUpstreamTotal) t.Errorf("log line has time %q, request_id %q and durations %v and %v",
line.Time, line.RequestID, line.DurationTotal,
line.fields["duration_upstream_total"])
} }
} }
@@ -371,8 +381,13 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
addr, out := startProxy(t, "http://"+localhost+":1", nil) addr, out := startProxy(t, "http://"+localhost+":1", nil)
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway) wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
wantLine(t, out.requestLine(t), http.StatusBadGateway,
requestlog.ActionUpstreamError) line := out.requestLine(t)
wantLine(t, line, http.StatusBadGateway, requestlog.ActionUpstreamError)
// There never was a connection to the app, nor an answer from it.
wantTimings(t, line, "duration_total", "duration_checks",
"duration_upstream_total")
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool { logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
return line["type"] == "process" && line["msg"] == "request to the app failed" return line["type"] == "process" && line["msg"] == "request to the app failed"
+5
View File
@@ -160,6 +160,9 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// a health checker is never refused. It does not ask the app. // a health checker is never refused. It does not ask the app.
if r.Method == http.MethodGet && r.URL.Path == HealthPath { if r.Method == http.MethodGet && r.URL.Path == HealthPath {
rq.line.Action = requestlog.ActionAdmin rq.line.Action = requestlog.ActionAdmin
// Set here rather than left to Go's server, which would set it only
// after the log line has taken the response's headers.
rq.out.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = io.WriteString(rq.out, "ok\n") _, _ = io.WriteString(rq.out, "ok\n")
return return
@@ -169,6 +172,8 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
defer rq.addToHistory() defer rq.addToHistory()
refused := rq.check(r.Context()) refused := rq.check(r.Context())
rq.checked = time.Now()
if refused != nil { if refused != nil {
rq.answer(*refused) rq.answer(*refused)
+15 -1
View File
@@ -35,6 +35,10 @@ const (
// localhost is where every test server listens, and so the address // localhost is where every test server listens, and so the address
// smallwebwaf sees each test's requests come from. // smallwebwaf sees each test's requests come from.
localhost = "127.0.0.1" localhost = "127.0.0.1"
// requestType is the type that marks a request log line.
requestType = "request"
// protocol is the protocol of every test's requests.
protocol = "HTTP/1.1"
) )
// shortTimeoutSetting is shortTimeout as a setting's value. // shortTimeoutSetting is shortTimeout as a setting's value.
@@ -68,6 +72,8 @@ const (
maxBanDuration = "SWWAF_MAX_BAN_DURATION" maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS" maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX" banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
) )
// output collects what smallwebwaf writes on stdout. // output collects what smallwebwaf writes on stdout.
@@ -84,6 +90,14 @@ func (o *output) Write(p []byte) (int, error) {
return o.buf.Write(p) return o.buf.Write(p)
} }
// text returns everything written so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// lines returns every line written so far, decoded. // lines returns every line written so far, decoded.
func (o *output) lines(t *testing.T) []map[string]any { func (o *output) lines(t *testing.T) []map[string]any {
t.Helper() t.Helper()
@@ -123,7 +137,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
var found []logLine var found []logLine
for _, fields := range o.lines(t) { for _, fields := range o.lines(t) {
if fields["type"] == "request" { if fields["type"] == requestType {
found = append(found, decodeLine(t, fields)) found = append(found, decodeLine(t, fields))
} }
} }
+23 -9
View File
@@ -6,6 +6,7 @@ import (
"testing" "testing"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -92,16 +93,14 @@ func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
// With a limit of one request a minute, the requests for paths under a // With a limit of one request a minute, the requests for paths under a
// prefix are not counted, so client's first request for / is within // prefix are not counted, so client's first request for / is within
// the limit; and once client has reached it, they are not refused. // the limit; and once client has reached it, they are not refused.
// The app acts on /static/../assets/app.js as /assets/app.js.
s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward) s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward) s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusOK, requestlog.ActionForward)
s.request(client, "/static/../assets/app.js",
http.StatusOK, requestlog.ActionForward)
line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward) line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
if line.LimitHit != "" { if line.LimitHit != "" || line.Counts != (ratelimit.Counts{}) {
t.Errorf("log line has limit_hit %q, want none", line.LimitHit) t.Errorf("log line has limit_hit %q and counts %+v, want neither",
line.LimitHit, line.Counts)
} }
// A path outside every prefix is counted: /assets is not under // A path outside every prefix is counted: /assets is not under
@@ -116,17 +115,32 @@ func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
http.StatusForbidden, requestlog.ActionCountryDenied) http.StatusForbidden, requestlog.ActionCountryDenied)
} }
func TestRateLimitCountsPathsOutsideEveryExemptPrefix(t *testing.T) { func TestRateLimitCountsPathsThatAreNotExempt(t *testing.T) {
t.Parallel() t.Parallel()
// A prefix matches only at the start of the path, and the app acts on
// the last three paths as /login, once they are percent-decoded and
// their .. segments resolved.
for _, sent := range []string{ for _, sent := range []string{
// A prefix matches only at the start of the path.
"/static/assets/app.js", "/static/assets/app.js",
// A prefix matches the path as sent: a router that matches the
// path as received does not take /%61ssets/x for a path under
// /assets/.
"/%61ssets/x",
// .. once percent-decoded: an app may act on these as /login, the
// last as a path under /sneak/app/ or as /assets/x.
"/assets/../login", "/assets/../login",
"/assets/%2e%2e/login", "/assets/%2e%2e/login",
"/assets/..%2Flogin", "/assets/..%2Flogin",
"/assets/..;/login",
"/sneak/app/src/branch/main/..%2F..%2F..%2F..%2F..%2F..%2Fassets/x",
// Not under /assets/ as sent: Go's router takes /assets%2Fx for one
// path segment, not a path under /assets/.
"/assets%2Fx",
"/assets%2fx",
// Under /assets/ as sent, but holding an encoded slash, in either
// case, or a backslash: never exempt, whatever the prefix.
"/assets/x%2Fy",
"/assets/x%2fy",
`/assets/x\y`,
} { } {
t.Run(sent, func(t *testing.T) { t.Run(sent, func(t *testing.T) {
t.Parallel() t.Parallel()
+137 -35
View File
@@ -7,8 +7,8 @@ import (
"net/http/httptrace" "net/http/httptrace"
"net/http/httputil" "net/http/httputil"
"net/netip" "net/netip"
"net/url"
"os" "os"
"path"
"slices" "slices"
"strings" "strings"
"sync" "sync"
@@ -49,7 +49,9 @@ type request struct {
peer netip.Addr peer netip.Addr
peerTrusted bool peerTrusted bool
start time.Time start time.Time
// upstreamStart is when the request was handed to the app. // checked is when the checks were done, and upstreamStart when the
// request was handed to the app.
checked time.Time
upstreamStart time.Time upstreamStart time.Time
// cancel ends the request to the app. // cancel ends the request to the app.
cancel context.CancelFunc cancel context.CancelFunc
@@ -59,26 +61,34 @@ type request struct {
complete bool complete bool
// mu guards what follows. The timeouts run on goroutines of their // mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them from its own; once // own, and the transport starts and stops them, and notes the times
// timersStopped is set, none of them acts any more. // below, from its own; once timersStopped is set, none of the timeouts
// acts any more.
mu sync.Mutex mu sync.Mutex
timersStopped bool timersStopped bool
clientRequestTimer *time.Timer clientRequestTimer *time.Timer
upstreamRequestTimer *time.Timer upstreamRequestTimer *time.Timer
upstreamResponseTimer *time.Timer upstreamResponseTimer *time.Timer
// requestSent is when the app had been sent the whole request. // connected is when there was a connection to the app, requestSent
requestSent time.Time // when the app had been sent the whole request, and answerStarted
// when the first byte of its answer arrived.
connected time.Time
requestSent time.Time
answerStarted time.Time
} }
// newRequest starts handling r: it notes the time, counts the request as // newRequest starts handling r: it notes the time, counts the request as
// under way, and works out the client. // under way, works out the client, and starts the log line with what is
// known of the request.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request { func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
h.metrics.RequestStarted() h.metrics.RequestStarted()
start := time.Now() start := time.Now()
peer := peerAddress(r) peer := peerAddress(r)
trusted := h.config.TrustedProxies trusted := h.config.TrustedProxies
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted) peerTrusted := isInside(peer, trusted)
forwardedFor := r.Header.Values("X-Forwarded-For")
client := clientAddress(peer, forwardedFor, trusted)
rq := &request{ rq := &request{
h: h, h: h,
@@ -87,22 +97,37 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
out: &responseWriter{ResponseWriter: w}, out: &responseWriter{ResponseWriter: w},
client: client, client: client,
peer: peer, peer: peer,
peerTrusted: isInside(peer, trusted), peerTrusted: peerTrusted,
start: start, start: start,
line: requestlog.Line{ line: requestlog.Line{
Time: requestlog.FormatTime(start), Time: requestlog.FormatTime(start),
ClientIP: client.String(), Instance: h.config.InstanceName,
PeerIP: peer.String(), ClientIP: client.String(),
Method: r.Method, Method: r.Method,
Host: r.Host, Scheme: scheme(r, peerTrusted),
Path: r.URL.EscapedPath(), Host: r.Host,
Query: r.URL.RawQuery, Path: r.URL.EscapedPath(),
Protocol: r.Proto, Query: r.URL.RawQuery,
Referer: r.Referer(), Protocol: r.Proto,
UserAgent: r.UserAgent(), Referer: r.Referer(),
Action: requestlog.ActionForward, UserAgent: r.UserAgent(),
RequestID: requestID(r, peerTrusted),
PeerIP: peer.String(),
ForwardedFor: strings.Join(forwardedFor, ", "),
ClientGroup: clientGroup(client).String(),
ContentType: r.Header.Get("Content-Type"),
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
HasAuthorization: len(r.Header.Values("Authorization")) > 0,
HasCookie: len(r.Header.Values("Cookie")) > 0,
Action: requestlog.ActionForward,
}, },
} }
// A length of -1 is a body whose length was not announced.
if r.ContentLength > 0 {
rq.line.ContentLength = r.ContentLength
}
if r.Body != http.NoBody { if r.Body != http.NoBody {
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq} rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
} }
@@ -110,6 +135,27 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
return rq return rq
} }
// requestHeaders returns the headers of r that names lists, by name in
// lower case, each with its values joined by ", ". Authorization, Cookie
// and Set-Cookie are never among them, whatever names says.
func requestHeaders(r *http.Request, names []string) map[string]string {
headers := map[string]string{}
for _, name := range names {
switch name {
case "authorization", "cookie", "set-cookie":
continue
}
values := r.Header.Values(name)
if len(values) > 0 {
headers[name] = strings.Join(values, ", ")
}
}
return headers
}
// check is the one place where a request can be refused once its client // check is the one place where a request can be refused once its client
// is known, before its body is read or anything reaches the app. It // is known, before its body is read or anything reaches the app. It
// returns nil to let the request through. The checks of checkClient come // returns nil to let the request through. The checks of checkClient come
@@ -148,9 +194,9 @@ func (rq *request) check(ctx context.Context) *refusal {
// client either refuses is not looked up, and then the country lists; a // client either refuses is not looked up, and then the country lists; a
// request any of them refuses is not counted for the rate limits. Then // request any of them refuses is not counted for the rate limits. Then
// come the rate limits, unless the client is in // come the rate limits, unless the client is in
// SWWAF_RATE_LIMIT_EXEMPT_NETS or the path the app will act on starts // SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under
// with one of SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request // SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted.
// is counted. ctx is the request's own context. // ctx is the request's own context.
func (rq *request) checkClient(ctx context.Context) string { func (rq *request) checkClient(ctx context.Context) string {
cfg := rq.h.config cfg := rq.h.config
if isInside(rq.client, cfg.AllowNets) { if isInside(rq.client, cfg.AllowNets) {
@@ -171,13 +217,8 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionCountryDenied return requestlog.ActionCountryDenied
} }
// The prefixes are matched against the path the app will act on:
// URL.Path is the request's path percent-decoded, and path.Clean
// resolves its . and .. segments and repeated slashes, and drops a
// trailing slash. So /assets/..%2Flogin is /login, outside /assets/,
// whatever the log line's path shows.
exempt := isInside(rq.client, cfg.RateLimitExemptNets) || exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
startsWithAny(path.Clean(rq.in.URL.Path), cfg.RateLimitExemptPaths) pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if !exempt && rq.limitBroken(now) { if !exempt && rq.limitBroken(now) {
return requestlog.ActionRateLimited return requestlog.ActionRateLimited
} }
@@ -185,10 +226,27 @@ func (rq *request) checkClient(ctx context.Context) string {
return "" return ""
} }
// startsWithAny reports whether s starts with one of prefixes. // pathExempt reports whether the rate limits leave out a request for u
func startsWithAny(s string, prefixes []string) bool { // because of SWWAF_RATE_LIMIT_EXEMPT_PATHS: whether its path as sent, the
// path the app receives, not percent-decoded, starts with one of
// prefixes, so that /%61ssets/x is not under /assets/ for an app whose
// router matches the path as received. A request whose decoded path
// contains .. anywhere or a backslash, or whose path as sent holds an
// encoded slash (%2F or %2f), never is, since an app may act on it as a
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
// as one path segment, as Go's router does.
func pathExempt(u *url.URL, prefixes []string) bool {
decoded := u.Path
// EscapedPath is the path as the app receives it, not decoded.
sent := u.EscapedPath()
if strings.Contains(decoded, "..") || strings.Contains(decoded, `\`) ||
strings.Contains(strings.ToLower(sent), "%2f") {
return false
}
return slices.ContainsFunc(prefixes, func(prefix string) bool { return slices.ContainsFunc(prefixes, func(prefix string) bool {
return strings.HasPrefix(s, prefix) return strings.HasPrefix(sent, prefix)
}) })
} }
@@ -200,7 +258,9 @@ func (rq *request) forward(ctx context.Context) {
rq.cancel = cancel rq.cancel = cancel
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
WroteRequest: rq.wroteRequest, GotConn: rq.gotConn,
WroteRequest: rq.wroteRequest,
GotFirstResponseByte: rq.gotFirstResponseByte,
}) })
out := rq.in.WithContext(ctx) out := rq.in.WithContext(ctx)
@@ -223,7 +283,8 @@ func (rq *request) forward(ctx context.Context) {
} }
// rewrite makes the request the app receives: the client's request, // rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set. // unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set.
func (rq *request) rewrite(pr *httputil.ProxyRequest) { func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme pr.Out.URL.Scheme = upstream.Scheme
@@ -232,6 +293,7 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
// the query as the client sent it. // the query as the client sent it.
pr.Out.URL.RawQuery = pr.In.URL.RawQuery pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
} }
// modifyResponse looks at the app's answer before ReverseProxy passes it // modifyResponse looks at the app's answer before ReverseProxy passes it
@@ -245,6 +307,7 @@ func (rq *request) modifyResponse(res *http.Response) error {
// connection it takes over, not through rq.out. // connection it takes over, not through rq.out.
rq.stopTimers() rq.stopTimers()
rq.out.status = res.StatusCode rq.out.status = res.StatusCode
rq.line.Websocket = true
return nil return nil
} }
@@ -340,6 +403,10 @@ func (rq *request) finish() {
line := &rq.line line := &rq.line
line.Status = rq.out.status line.Status = rq.out.status
line.ResponseBytes = rq.out.bytes line.ResponseBytes = rq.out.bytes
header := rq.out.Header()
line.ResponseContentType = header.Get("Content-Type")
line.CacheControl = header.Get("Cache-Control")
line.Location = header.Get("Location")
if rq.body != nil { if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load() line.RequestBytes = rq.body.bytes.Load()
@@ -364,12 +431,18 @@ func (rq *request) finish() {
now := time.Now() now := time.Now()
duration := now.Sub(rq.start) duration := now.Sub(rq.start)
line.DurationTotal = requestlog.Milliseconds(duration) line.DurationTotal = requestlog.Milliseconds(duration)
line.DurationChecks = timing(rq.start, rq.checked)
var upstreamDuration time.Duration var upstreamDuration time.Duration
if !rq.upstreamStart.IsZero() { if !rq.upstreamStart.IsZero() {
upstreamDuration = now.Sub(rq.upstreamStart) upstreamDuration = now.Sub(rq.upstreamStart)
line.DurationUpstreamTotal = requestlog.Milliseconds(upstreamDuration) line.DurationUpstreamTotal = new(requestlog.Milliseconds(upstreamDuration))
rq.mu.Lock()
line.DurationUpstreamConnect = timing(rq.upstreamStart, rq.connected)
line.DurationUpstreamFirstByte = timing(rq.upstreamStart, rq.answerStarted)
rq.mu.Unlock()
} }
// Counted before the log line is written, so that the metrics count // Counted before the log line is written, so that the metrics count
@@ -382,6 +455,17 @@ func (rq *request) finish() {
} }
} }
// timing is the time from start to end in milliseconds, for one of the
// log line's timings, or nil when end is zero: what it times never
// happened.
func timing(start, end time.Time) *float64 {
if end.IsZero() {
return nil
}
return new(requestlog.Milliseconds(end.Sub(start)))
}
// addToHistory adds the request, which has ended, to its client's // addToHistory adds the request, which has ended, to its client's
// history. // history.
func (rq *request) addToHistory() { func (rq *request) addToHistory() {
@@ -491,6 +575,24 @@ func (rq *request) bodyReceived() {
stopTimer(rq.clientRequestTimer) stopTimer(rq.clientRequestTimer)
} }
// gotConn is called once there is a connection to the app, a new one or
// one kept open from an earlier request.
func (rq *request) gotConn(httptrace.GotConnInfo) {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.connected = time.Now()
}
// gotFirstResponseByte is called once the first byte of the app's answer
// has arrived.
func (rq *request) gotFirstResponseByte() {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.answerStarted = time.Now()
}
// wroteRequest is called once the app has been sent the whole request: // wroteRequest is called once the app has been sent the whole request:
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts. // the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) { func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
+368
View File
@@ -0,0 +1,368 @@
package proxy_test
import (
"io"
"maps"
"math"
"net/http"
"reflect"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// requestIDHeader carries the request's id.
requestIDHeader = "X-Request-ID"
// instance is the SWWAF_INSTANCE_NAME a test sets.
instance = "fsn1app1/gitea"
// ipv6Client is a client on IPv6, and ipv6Group the netblock the rate
// limits count it as.
ipv6Client = "2001:db8::7"
ipv6Group = "2001:db8::/64"
)
func TestLogLineHasEachFieldWhereItApplies(t *testing.T) {
t.Parallel()
received := make(chan string, 2) // the request ids the app received
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path != "/full" {
w.WriteHeader(http.StatusNoContent)
return
}
w.Header().Set("Content-Type", "text/html")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Location", "/elsewhere")
w.WriteHeader(http.StatusFound)
_, _ = io.WriteString(w, "moved")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
rateLimitExemptNets: localhost,
instanceName: instance,
logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE",
})
// This request comes from ipv6Client through a trusted proxy, with a
// body and each header the log line looks at, and is answered with a
// redirect.
conn := dial(t, addr)
send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+
forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+
forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+
"Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+
"Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+
"Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+
"Set-Cookie: secret-set-cookie\r\n\r\na=b")
wantStatus(t, readResponse(t, conn), http.StatusFound)
// A request's log line can come after its answer: each is waited for
// before the next request, so that the lines are in order.
full := out.requestLines(t, 1)[0]
// This one comes from 127.0.0.1, which the rate limits do not count,
// with a body of 4 bytes whose length it does not announce, so that its
// request_bytes is not its content_length, and no header the log line
// looks at, and is answered with 204 and no header.
conn = dial(t, addr)
send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+
"Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusNoContent)
bare := out.requestLines(t, 2)[1]
wantFullLine(t, full)
wantBareLine(t, bare)
for _, line := range []logLine{full, bare} {
got := <-received
if got != line.RequestID {
t.Errorf("the app received request id %q, the log line has %q",
got, line.RequestID)
}
}
if strings.Contains(out.text(), "secret") {
t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s",
out.text())
}
}
// wantFullLine checks the log line of the request with every header the
// line looks at. Its timings are checked by TestTimingsAreInOrder.
func wantFullLine(t *testing.T, line logLine) {
t.Helper()
headers := map[string]string{"accept": "text/html", "x-custom": "one, two"}
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure,
Host: appHost, Path: "/full", Protocol: protocol,
Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5,
RequestID: "from-traefik", PeerIP: localhost,
ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group,
ContentType: "application/x-www-form-urlencoded", ContentLength: 3,
RequestHeaders: headers, HasAuthorization: true, HasCookie: true,
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
}
// wantBareLine checks the log line of the request with none of them, and
// that the fields that do not apply to it are left out.
func wantBareLine(t *testing.T, line logLine) {
t.Helper()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: localhost, Method: http.MethodPost, Scheme: plain,
Host: appHost, Path: "/bare", Protocol: protocol,
Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID,
PeerIP: localhost, ClientGroup: localhost + "/32",
UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward,
})
if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" {
t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want)
}
for _, name := range []string{
"forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"cache_control", "location", "counts",
} {
_, present := line.fields[name]
if present {
t.Errorf("log line has %s, which does not apply", name)
}
}
}
// withTimings returns want with the timings of line.
func withTimings(line logLine, want requestlog.Line) requestlog.Line {
want.DurationTotal = line.DurationTotal
want.DurationChecks = line.DurationChecks
want.DurationUpstreamConnect = line.DurationUpstreamConnect
want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte
want.DurationUpstreamTotal = line.DurationUpstreamTotal
return want
}
func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) {
t.Parallel()
const hasAuthorization, hasCookie = "has_authorization", "has_cookie"
for _, tc := range []struct{ header, field, other string }{
{"Authorization", hasAuthorization, hasCookie},
{"Cookie", hasCookie, hasAuthorization},
} {
t.Run("only "+tc.header, func(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out := startProxy(t, app.URL, nil)
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(tc.header, "secret")
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLine(t)
_, otherPresent := line.fields[tc.other]
if line.fields[tc.field] != true || otherPresent {
t.Errorf("log line has %s %v and %s %v, want true and none",
tc.field, line.fields[tc.field], tc.other, line.fields[tc.other])
}
})
}
}
func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) {
t.Parallel()
const sentID = "from-traefik"
sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}}
trusted := map[string]string{trustedProxies: trustLocalhost}
for _, tc := range []struct {
name string
env map[string]string
header http.Header
// wantID is the request id logged, "" for a new one.
wantID, wantScheme string
}{
{"a trusted proxy's are kept", trusted, sent, sentID, secure},
{"without them, the id is new and the scheme http", trusted, nil, "", plain},
{"another peer's are replaced", nil, sent, "", plain},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
received := make(chan string, 2)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
})
addr, out := startProxy(t, app.URL, tc.env)
// Two requests, so that two new ids can be told apart.
ids := make([]string, 0, 2)
for i := range 2 {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
maps.Copy(req.Header, tc.header)
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLines(t, i+1)[i]
ids = append(ids, line.RequestID)
got := <-received
if line.RequestID != got || line.Scheme != tc.wantScheme {
t.Errorf("log line has request_id %q and scheme %q, and the "+
"app received id %q; want the same id and scheme %q",
line.RequestID, line.Scheme, got, tc.wantScheme)
}
}
switch {
case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID):
t.Errorf("request ids %q, want %q", ids, tc.wantID)
case tc.wantID == "" && (slices.Contains(ids, sentID) ||
slices.Contains(ids, "") || ids[0] == ids[1]):
t.Errorf("request ids %q, want two new ones", ids)
}
})
}
}
func TestTimingsAreInOrder(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
// The pauses set the times apart; a hold-up of the test only
// lengthens them.
time.Sleep(time.Millisecond)
w.WriteHeader(http.StatusOK)
_ = http.NewResponseController(w).Flush()
time.Sleep(time.Millisecond)
_, _ = io.WriteString(w, "done")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
denyNets: denied,
})
// Each log line is waited for before the next request, so that the
// lines are in order.
wantStatus(t, get(t, addr, "/"), http.StatusOK)
forwarded := out.requestLines(t, 1)[0]
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, denied)
wantStatus(t, do(t, req), http.StatusForbidden)
refused := out.requestLines(t, 2)[1]
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
health := out.requestLines(t, 3)[2]
// A request passed to the app has every timing; one refused, none of
// the app's; the health check, which runs no check, only the total.
wantTimings(t, forwarded, "duration_total", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total")
wantTimings(t, refused, "duration_total", "duration_checks")
wantTimings(t, health, "duration_total")
if t.Failed() {
return
}
// In whole microseconds, as they are logged, so that the sum below is
// exact.
total := microseconds(forwarded.DurationTotal)
checks := microseconds(*forwarded.DurationChecks)
connect := microseconds(*forwarded.DurationUpstreamConnect)
firstByte := microseconds(*forwarded.DurationUpstreamFirstByte)
upstream := microseconds(*forwarded.DurationUpstreamTotal)
// The checks end before the request is handed to the app, and the
// connection comes before the answer, which the app ends after a
// pause.
if checks+upstream > total || connect >= firstByte || firstByte >= upstream {
t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+
"first byte %d, upstream total %d", total, checks, connect, firstByte,
upstream)
}
if *refused.DurationChecks > refused.DurationTotal {
t.Errorf("refused request's checks took %v of %v milliseconds",
*refused.DurationChecks, refused.DurationTotal)
}
}
// wantTimings checks that the timings named are the only ones line has.
func wantTimings(t *testing.T, line logLine, want ...string) {
t.Helper()
var got []string
for name := range line.fields {
if strings.HasPrefix(name, "duration_") {
got = append(got, name)
}
}
slices.Sort(got)
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want)
}
}
// microseconds is a timing in whole microseconds.
func microseconds(milliseconds float64) int64 {
return int64(math.Round(milliseconds * 1000))
}
func TestLogsAnUpgradedConnection(t *testing.T) {
t.Parallel()
app := startApp(t, echoAfterUpgrade)
addr, out := startProxy(t, app.URL, nil)
conn := dial(t, addr)
send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols)
_ = conn.Close()
line := out.requestLine(t)
if line.fields["websocket"] != true {
t.Errorf("log line has websocket %v, want true", line.fields["websocket"])
}
}
+31 -15
View File
@@ -149,26 +149,39 @@ type Hit struct {
Requests float64 Requests float64
} }
// Counts are a client's requests in the minute, the hour and the day that
// end at a request, that request included.
type Counts struct {
Minute float64 `json:"minute"`
Hour float64 `json:"hour"`
Day float64 `json:"day"`
}
// Count counts a request from client at now, in every window, whether or // Count counts a request from client at now, in every window, whether or
// not it is refused. It reports whether the request takes the client over // not it is refused, and returns the client's requests in each window. It
// a limit, and the window whose limit it goes over, the shortest if it is // reports whether the request takes the client over a limit, and the
// over several. // window whose limit it goes over, the shortest if it is over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) { func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
var hit Hit var (
requests [3]float64
hit Hit
)
for i, b := range l.get(client).buckets() { for i, b := range l.get(client).buckets() {
w := l.windows[i] w := l.windows[i]
requests := b.add(now, w.length) requests[i] = b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) { if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests} hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
} }
} }
return hit, hit.Window != "" counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
return counts, hit, hit.Window != ""
} }
// Reset sets client's counts in every window back to zero. Its history // Reset sets client's counts in every window back to zero. Its history
@@ -269,18 +282,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
+26 -3
View File
@@ -62,14 +62,14 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight() start := midnight()
for range limit { for range limit {
_, over := limiter.Count(client, start) _, _, over := limiter.Count(client, start)
if over { if over {
t.Fatal("a request within the limit is over it") t.Fatal("a request within the limit is over it")
} }
} }
// Over both limits; the minute's is named, with the four requests. // Over both limits; the minute's is named, with the four requests.
hit, over := limiter.Count(client, start) _, hit, over := limiter.Count(client, start)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1} want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want { if !over || hit != want {
@@ -78,6 +78,29 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
} }
} }
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range 3 {
limiter.Count(client, start)
}
// A quarter into the next hour, the minute has only this request. The
// hour still covers three quarters of the bucket before, with its three
// requests, which count 2.25, and this one: 3.25. The day covers all
// four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want {
t.Errorf("counts %+v, want %+v", counts, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) { func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel() t.Parallel()
@@ -238,7 +261,7 @@ func wantCount(
) { ) {
t.Helper() t.Helper()
hit, _ := limiter.Count(client, now) _, hit, _ := limiter.Count(client, now)
if hit.Window != want { if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q", t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want) client, now.Format(time.RFC3339), hit.Window, want)
+72 -24
View File
@@ -9,6 +9,8 @@ import (
"io" "io"
"log/slog" "log/slog"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
) )
// The action a request line names: what smallwebwaf did with the // The action a request line names: what smallwebwaf did with the
@@ -45,32 +47,71 @@ const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds. // timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00" const timeLayout = "2006-01-02T15:04:05.000Z07:00"
// Line is one request's line in the request log. The field names are // Line is one request's line in the request log. The field names, and
// those of the "Request log" section of SPEC.md. // their order, are those of the "Request log" section of SPEC.md. A field
// that may not apply to a request is left out of its line when it does
// not.
// //
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case //nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Line struct { type Line struct {
Type string `json:"type"` Type string `json:"type"`
Time string `json:"time"`
ClientIP string `json:"client_ip"` // The standard web log fields. Scheme is how the client reached
PeerIP string `json:"peer_ip"` // smallwebwaf, or the trusted proxy in front of it.
Country string `json:"country"` Time string `json:"time"`
Method string `json:"method"` Instance string `json:"instance"`
Host string `json:"host"` ClientIP string `json:"client_ip"`
Path string `json:"path"` Method string `json:"method"`
Query string `json:"query"` Scheme string `json:"scheme"`
Protocol string `json:"protocol"` Host string `json:"host"`
Status int `json:"status"` Path string `json:"path"`
UpstreamStatus int `json:"upstream_status,omitempty"` Query string `json:"query"`
RequestBytes int64 `json:"request_bytes"` Protocol string `json:"protocol"`
ResponseBytes int64 `json:"response_bytes"` Status int `json:"status"`
Referer string `json:"referer"` RequestBytes int64 `json:"request_bytes"`
UserAgent string `json:"user_agent"` ResponseBytes int64 `json:"response_bytes"`
Action string `json:"action"` Referer string `json:"referer"`
UserAgent string `json:"user_agent"`
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
// or a new one, and is sent on to the app. ForwardedFor is the
// X-Forwarded-For header as received. ClientGroup is the netblock the
// client is counted as.
RequestID string `json:"request_id"`
PeerIP string `json:"peer_ip"`
ForwardedFor string `json:"forwarded_for,omitempty"`
ClientGroup string `json:"client_group"`
Country string `json:"country"`
ContentType string `json:"content_type,omitempty"`
// ContentLength is the length of its body the request announced.
ContentLength int64 `json:"content_length,omitempty"`
// RequestHeaders are the headers SWWAF_LOG_REQUEST_HEADERS names that
// the request carried, by name in lower case.
RequestHeaders map[string]string `json:"request_headers,omitempty"`
HasAuthorization bool `json:"has_authorization,omitempty"`
HasCookie bool `json:"has_cookie,omitempty"`
// Websocket is true when the connection was upgraded, as for a
// WebSocket.
Websocket bool `json:"websocket,omitempty"`
// Response detail, from the headers of the answer: the app's, as
// passed on, or those of smallwebwaf's own. Aborted is true when the
// client went away early.
ResponseContentType string `json:"response_content_type,omitempty"`
UpstreamStatus int `json:"upstream_status,omitempty"`
CacheControl string `json:"cache_control,omitempty"`
Location string `json:"location,omitempty"`
Aborted bool `json:"aborted,omitempty"`
// The decision.
Action string `json:"action"`
// WouldAction is, in observe mode, the action enforce mode would have // WouldAction is, in observe mode, the action enforce mode would have
// taken with a request it would have refused: ActionDenied, // taken with a request it would have refused: ActionDenied,
// ActionBanned, ActionCountryDenied or ActionRateLimited. // ActionBanned, ActionCountryDenied or ActionRateLimited.
WouldAction string `json:"would_action,omitempty"` WouldAction string `json:"would_action,omitempty"`
// Counts are the client's requests as the rate limits counted them
// with this one, for a request they counted.
Counts ratelimit.Counts `json:"counts,omitzero"`
// LimitHit is the window whose rate limit the request went over: // LimitHit is the window whose rate limit the request went over:
// minute, hour or day. // minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"` LimitHit string `json:"limit_hit,omitempty"`
@@ -79,11 +120,18 @@ type Line struct {
// BanExpires is when the ban the request made, or was refused under, // BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent". // ends: a time, or "permanent".
BanExpires string `json:"ban_expires,omitempty"` BanExpires string `json:"ban_expires,omitempty"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"` // The timings, in milliseconds. DurationChecks is the time until the
// DurationTotal and DurationUpstreamTotal are in milliseconds. // checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte
DurationTotal float64 `json:"duration_total"` // and DurationUpstreamTotal run from when the request was handed to the
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"` // app: until there was a connection to it, until the first byte of its
// answer arrived, and until the end. Each but DurationTotal is nil for
// a request that did not get that far.
DurationTotal float64 `json:"duration_total"`
DurationChecks *float64 `json:"duration_checks,omitempty"`
DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"`
DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"`
DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"`
} }
// Write writes line to w as one JSON line marked "type":"request". // Write writes line to w as one JSON line marked "type":"request".
+5 -1
View File
@@ -50,7 +50,11 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
} }
unset := []string{ unset := []string{
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted", "forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"upstream_status", "cache_control", "location", "aborted", "counts",
"limit_hit", "offence", "ban_expires", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total", "duration_upstream_total",
} }
for _, name := range unset { for _, name := range unset {
+18 -9
View File
@@ -113,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,
@@ -130,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())
@@ -165,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 {
+86 -19
View File
@@ -26,11 +26,14 @@ const (
// testVersion is the version the tests give smallwebwaf. // testVersion is the version the tests give smallwebwaf.
testVersion = "test" testVersion = "test"
// localhost is where the tests listen. // localhost is where the tests listen.
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"
stateDir = "SWWAF_STATE_DIR" trustedProxies = "SWWAF_TRUSTED_PROXIES"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
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"
) )
@@ -217,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.
@@ -247,12 +250,12 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
const scope = "SWWAF_BAN_SCOPE_V4_PREFIX" const scope = "SWWAF_BAN_SCOPE_V4_PREFIX"
env := map[string]string{ env := map[string]string{
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",
} }
// 203.0.113.9's second request breaks the day limit, and bans // 203.0.113.9's second request breaks the day limit, and bans
@@ -281,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()
@@ -384,9 +419,9 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
upstreamURL: appURL, upstreamURL: appURL,
stateDir: dir, stateDir: dir,
"SWWAF_MODE": "enforce", "SWWAF_MODE": "enforce",
"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",
@@ -480,6 +515,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 {
@@ -498,7 +567,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)
}
} }
+270 -76
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,8 +20,10 @@ 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/metrics"
@@ -61,15 +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. // Metrics count each file's writes, and the edits taken in and set
// aside.
Metrics *metrics.Metrics 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.
@@ -120,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()
@@ -172,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())
} }
@@ -182,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.fileChanged(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.
@@ -193,52 +239,209 @@ func (f *Files) logFailure(err error) {
} }
} }
// writeBans writes bans.json. // fileChanged takes in what the state file name holds, as Watch sees it
func (f *Files) writeBans() error { // change, if that is an edit made since smallwebwaf last read or wrote
held := f.params.Ledger.Snapshot() // the file. A file that cannot be read or does not parse is left for its
// next write.
func (f *Files) fileChanged(name string) {
f.mu.Lock()
defer f.mu.Unlock()
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))} data, changed, err := f.readChanged(name)
for _, ban := range held { if err != nil || !changed {
file.Bans = append(file.Bans, newBanEntry(ban)) return
} }
data, err := json.MarshalIndent(file, "", " ") _ = f.takeInEdit(name, data)
if err != nil {
return fmt.Errorf("encode %s: %w", bansJSON, err)
}
return f.writeCounted(bansJSON, append(data, '\n'))
} }
// writeClients writes clients.json. // takeInEdit takes in data, an edit of the state file name, as takeIn
func (f *Files) writeClients() error { // does, and counts and logs it. Every edit taken in while smallwebwaf
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot()) // runs, by Watch or by a write, is taken in here. An edit that does not
// parse is neither counted nor logged, and takeIn's error returned.
func (f *Files) takeInEdit(name string, data []byte) error {
_, err := f.takeIn(name, data)
if err != nil { if err != nil {
return fmt.Errorf("encode %s: %w", clientsJSON, err) return err
} }
return f.writeCounted(clientsJSON, data) // Counted before it is logged, so that the count is there once the
// log line is.
f.params.Metrics.StateFileEditTakenIn(name)
f.params.ProcessLog.Info("took in an edit of a state file",
"file", filepath.Join(f.params.Dir, name))
return nil
} }
// writeLookups writes lookups.json. // read takes in the state file name at start, and returns how many
func (f *Files) writeLookups() error { // entries it holds. A missing file holds none.
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) func (f *Files) read(name string) (int, error) {
if err != nil { data, changed, err := f.readChanged(name)
return fmt.Errorf("encode %s: %w", lookupsJSON, err) if err != nil || !changed {
return 0, err
} }
return f.writeCounted(lookupsJSON, data) return f.takeIn(name, data)
} }
// writeCounted writes data to the state file name, as write does, and // readChanged returns what the state file name holds, and whether that
// counts the write in the metrics. // has changed since smallwebwaf last read or wrote the file, as it has
func (f *Files) writeCounted(name string, data []byte) error { // for a file smallwebwaf never read or wrote. A missing file has not
err := write(f.params.Dir, name, data) // 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, or set aside if it does not
// parse. A file that cannot be read, or an edit that cannot be set
// aside, is left as it is, and the write given up. Every write is counted
// in the metrics, and one that fails or is given up as a failure.
func (f *Files) writeFile(name string) error {
f.mu.Lock()
defer f.mu.Unlock()
data, changed, err := f.readChanged(name)
if err == nil && changed {
err = f.takeInEdit(name, data)
if err != nil {
err = f.setAside(name, err)
}
}
if err == nil {
data, err = f.encode(name)
if err != nil {
err = fmt.Errorf("encode %s: %w", name, err)
}
}
if err == nil {
err = write(f.params.Dir, name, data)
}
if err == nil {
// The file holds data from here on, even if the directory sync
// fails, so that its next read does not take it for an admin's
// edit.
f.sums[name] = sha256.Sum256(data)
err = syncDirectory(f.params.Dir)
}
f.params.Metrics.StateFileWritten(name, len(data), err) f.params.Metrics.StateFileWritten(name, len(data), err)
return err return err
} }
// setAside renames the state file name, an edit that does not parse with
// parseErr, to name.bad, for the admin to mend, and logs it with where in
// the file the error is. If the rename fails, the edit is left as it is,
// and the error returned is parseErr joined with the rename's.
func (f *Files) setAside(name string, parseErr error) error {
path := filepath.Join(f.params.Dir, name)
err := os.Rename(path, path+".bad")
if err != nil {
return errors.Join(parseErr, err)
}
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse",
"file", path+".bad", "error", parseErr.Error())
f.params.Metrics.StateFileEditSetAside(name)
return nil
}
// encode returns the state file name as smallwebwaf writes it, from a
// snapshot of the part that keeps that state.
func (f *Files) encode(name string) ([]byte, error) {
switch name {
case bansJSON:
held := f.params.Ledger.Snapshot()
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))}
for _, ban := range held {
file.Bans = append(file.Bans, newBanEntry(ban))
}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
default: // lookups.json
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
}
}
// newBanEntry returns ban as bans.json holds it. // newBanEntry returns ban as bans.json holds it.
func newBanEntry(ban bans.Ban) banEntry { func newBanEntry(ban bans.Ban) banEntry {
entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes} entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes}
@@ -389,28 +592,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)
@@ -463,7 +654,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"
@@ -475,10 +666,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
+481 -11
View File
@@ -3,7 +3,10 @@ package state_test
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"io/fs"
"log/slog" "log/slog"
"maps"
"net"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
@@ -28,6 +31,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.
@@ -290,7 +300,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.
@@ -302,7 +312,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.
@@ -347,7 +357,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.
@@ -364,6 +374,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()
@@ -459,24 +507,359 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) {
float64(len(permanentBansJSON))) float64(len(permanentBansJSON)))
} }
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Parallel() t.Parallel()
dir := t.TempDir() dir := t.TempDir()
files := load(t, newParams(dir)) path := filepath.Join(dir, bansJSON)
params := newParams(dir)
files := load(t, params)
// A directory named bans.json cannot be renamed over. // bans.json is a socket, which cannot be opened as a file, even by
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) // 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)
wantWriteFailed(t, params, bansJSON)
}
func TestBrokenEditThatCannotBeSetAsideIsNotWrittenOver(t *testing.T) {
t.Parallel()
const broken = `{"version": 1, "bans": [`
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
files := load(t, params)
// A directory named bans.json.bad cannot be renamed over, so the
// broken edit cannot be set aside, and is left as it is.
edit(t, dir, bansJSON, broken)
err := os.Mkdir(path+".bad", 0o700)
if err != nil { if err != nil {
t.Fatalf("mkdir: %v", err) t.Fatalf("mkdir: %v", err)
} }
err = files.WriteAll() err = files.WriteAll()
if err == nil { if err == nil {
t.Error("writing over a directory did not fail") t.Error("writing with bans.json.bad in the way did not fail")
} }
if got := readFile(t, path); got != broken {
t.Errorf("bans.json holds\n%s\nwant the edit", got)
}
wantWriteFailed(t, params, bansJSON)
}
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 TestFileRenamedOverAStateFileTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// The admin mends bans.json.bad and moves it back, as editors that
// save by renaming do with a file of their own: nothing is written
// into bans.json itself. An edit of clients.json after it must be
// taken in second.
edit(t, dir, bansJSON+".bad", permanentBansJSON)
err = os.Rename(path+".bad", path)
if err != nil {
t.Fatalf("rename: %v", err)
}
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, bansJSON)
wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()})
}
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) 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 TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
// An edit taken in by Watch, which is then stopped.
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Watch(ctx)
close(stopped)
}()
lines.waitFor(t, watching)
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
byWatch := lines.waitFor(t, tookIn)
stop()
<-stopped
// An edit taken in by the write of its file. Nothing logs after the
// write, so the log is closed, and a write that does not log the edit
// fails the test at once instead of waiting for the line.
edit(t, dir, bansJSON, permanentBansJSON)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
close(lines)
byWrite := lines.waitFor(t, tookIn)
// The two lines differ only in their time.
delete(byWatch, "time")
delete(byWrite, "time")
if !maps.Equal(byWrite, byWatch) {
t.Errorf("the write logged %v, where Watch logged %v", byWrite, byWatch)
}
}
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.
@@ -573,15 +956,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)
}() }()
@@ -591,6 +975,80 @@ 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, or nil if the log is closed first. 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
}
// 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) {
@@ -735,6 +1193,18 @@ func metric(t *testing.T, text, series string) float64 {
return 0 return 0
} }
// wantWriteFailed checks that the metrics of params count one write of the
// state file name, and that it failed.
func wantWriteFailed(t *testing.T, params state.Params, name string) {
t.Helper()
got := scrape(t, params)
file := `{file="` + name + `"}`
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1)
}
// wantMetric checks the value of series in text, the metrics, as metric // wantMetric checks the value of series in text, the metrics, as metric
// reads it. // reads it.
func wantMetric(t *testing.T, text, series string, want float64) { func wantMetric(t *testing.T, text, series string, want float64) {
+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)
}
}