Compare commits
3
Commits
890dcedfd7
..
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
74bdc6a449 | ||
|
|
808e69f442 | ||
|
|
6ec52e5b87 |
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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) {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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"),
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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) {
|
||||||
|
|||||||
@@ -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"])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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".
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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(¶ms)
|
||||||
|
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) {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user