Keep the bans, the clients and GeoJS's answers in state files (closes #17)
check / check (push) Successful in 3m24s
check / check (push) Successful in 3m24s
smallwebwaf now copies its state to bans.json, clients.json and lookups.json in SWWAF_STATE_DIR, as "Persistent state" in SPEC.md describes, and reads them back at start, so a restart lifts no ban and gives no client a fresh allowance. Each client gains a history, and a ban's notes count the netblock's requests. bans.json is written SWWAF_STATE_WRITE_DELAY after a ban, and every file every SWWAF_STATE_COUNTER_INTERVAL and at the stop. A ban read back is masked to its netblock and refuses every client in it. A file that does not parse, an unknown version, an entry without a field it needs, or an unwritable directory stops the start. Deviation: no AS number or name, and no ban cause, reason or lifting yet. Model: opus-5-5
This commit was merged in pull request #72.
This commit is contained in:
@@ -162,6 +162,11 @@ RUN groupadd --system --gid 65532 smallwebwaf \
|
|||||||
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
|
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
|
||||||
smallwebwaf
|
smallwebwaf
|
||||||
|
|
||||||
|
# The state files' directory, SWWAF_STATE_DIR by default, where a volume
|
||||||
|
# is mounted to keep them across deploys. The run script gives it to the
|
||||||
|
# smallwebwaf user at each start.
|
||||||
|
RUN mkdir /var/lib/smallwebwaf
|
||||||
|
|
||||||
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
|
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
|
||||||
# looks too.
|
# looks too.
|
||||||
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
|
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
|
||||||
|
|||||||
@@ -13,18 +13,19 @@ 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 three parts of
|
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are four parts of
|
||||||
milestone 3: the static lists and the bans that broken rate limits lead to,
|
milestone 3: the static lists, the bans that broken rate limits lead to and the
|
||||||
which come next in the build order, and the header size and the idle time as
|
JSON state files, which come next in the build order, and the header size and
|
||||||
settings, which come last in it. `smallwebwaf` passes each request to the app
|
the idle time as settings, which come last in it. `smallwebwaf` passes each
|
||||||
and the app's answer back, unchanged, within its timeouts and size limits, works
|
request to the app and the app's answer back, unchanged, within its timeouts and
|
||||||
out each client's address, bans a client that sends too many requests, refuses a
|
size limits, works out each client's address, bans a client that sends too many
|
||||||
client that comes from a country you refuse or from a network you refuse, lets
|
requests, refuses a client that comes from a country you refuse or from a
|
||||||
the networks you choose through, and writes a JSON log line for every request.
|
network you refuse, lets the networks you choose through, keeps its bans, each
|
||||||
It comes as the image the app's own image is built on. The rest of the design
|
client's counters and history, and GeoJS's answers in JSON files across
|
||||||
comes after that, in the order of the build order in [`SPEC.md`](SPEC.md). The
|
restarts, and writes a JSON log line for every request. It comes as the image
|
||||||
survey of existing tools that led to the design is in
|
the app's own image is built on. The rest of the design comes after that, in the
|
||||||
[`EVALUATION.md`](EVALUATION.md).
|
order of the build order in [`SPEC.md`](SPEC.md). The survey of existing tools
|
||||||
|
that led to the design is in [`EVALUATION.md`](EVALUATION.md).
|
||||||
|
|
||||||
## Getting started
|
## Getting started
|
||||||
|
|
||||||
@@ -46,7 +47,8 @@ works.
|
|||||||
|
|
||||||
To work on the code, `make build` builds the binary alone, with Go installed,
|
To work on the code, `make build` builds the binary alone, with Go installed,
|
||||||
and `make run` builds and runs it, listening on port 8080 in front of an app at
|
and `make run` builds and runs it, listening on port 8080 in front of an app at
|
||||||
`SWWAF_UPSTREAM_URL`, by default `http://127.0.0.1:8081`.
|
`SWWAF_UPSTREAM_URL`, by default `http://127.0.0.1:8081`, with its state files
|
||||||
|
in `bin/state` unless `SWWAF_STATE_DIR` is set.
|
||||||
|
|
||||||
## What it does so far
|
## What it does so far
|
||||||
|
|
||||||
@@ -77,8 +79,9 @@ and `make run` builds and runs it, listening on port 8080 in front of an app at
|
|||||||
bans the client. A client is one IPv4 address, or one IPv6 /64, since one
|
bans the client. A client is one IPv4 address, or one IPv6 /64, since one
|
||||||
abuser usually holds a whole /64. Each window is counted in two fixed buckets,
|
abuser usually holds a whole /64. Each window is counted in two fixed buckets,
|
||||||
the earlier one weighted by how much of it the window still covers. At most
|
the earlier one weighted by how much of it the window still covers. At most
|
||||||
20,000 clients are kept, the least recently seen dropped first, and only in
|
20,000 clients are kept, the least recently seen dropped first, with their
|
||||||
memory: a restart starts every client afresh.
|
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
|
||||||
@@ -90,13 +93,13 @@ and `make run` builds and runs it, listening on port 8080 in front of an app at
|
|||||||
is not counted for the rate limits. A ban sets the client's counters back to
|
is not counted for the rate limits. A ban sets the client's counters back to
|
||||||
zero. Each ban carries notes for deciding whether to lift it: the limit, its
|
zero. Each ban carries notes for deciding whether to lift it: the limit, its
|
||||||
window and the requests counted in it, the request that broke it, the client's
|
window and the requests counted in it, the request that broke it, the client's
|
||||||
country when it was looked up, how many requests the ban has refused, and how
|
country when it was looked up, the netblock's requests since it was first
|
||||||
many bans the netblock had before. At most `SWWAF_MAX_BANS` bans are kept,
|
seen, how many of them the ban has refused, and how many bans the netblock had
|
||||||
past, active and permanent; past that, the earliest ban of the netblock that
|
before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent;
|
||||||
has gone longest without a request is dropped first. Bans and their notes are
|
past that, the earliest ban of the netblock that has gone longest without a
|
||||||
kept in memory only, so a restart lifts every ban, and nothing shows them yet:
|
request is dropped first. `bans.json` shows the bans and their notes, and a
|
||||||
`bans.json`, which shows them and lets you lift a ban, comes with the state
|
restart lifts none (see "State files" below); lifting a ban by editing it
|
||||||
files (https://git.eeqj.de/sneak/smallwebwaf/issues/17).
|
comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
|
||||||
- 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
|
||||||
@@ -184,6 +187,13 @@ it, and the effective settings are logged at start.
|
|||||||
- `SWWAF_BAN_SCOPE_V4_PREFIX` (default `32`): the length of the netblock around
|
- `SWWAF_BAN_SCOPE_V4_PREFIX` (default `32`): the length of the netblock around
|
||||||
an IPv4 client that a ban covers, such as `24` to ban the surrounding /24. An
|
an IPv4 client that a ban covers, such as `24` to ban the surrounding /24. An
|
||||||
IPv6 ban covers the client's /64.
|
IPv6 ban covers the client's /64.
|
||||||
|
- `SWWAF_STATE_DIR` (default `/var/lib/smallwebwaf`): the directory of the state
|
||||||
|
files, an absolute path. A directory `smallwebwaf` cannot write stops the
|
||||||
|
start.
|
||||||
|
- `SWWAF_STATE_WRITE_DELAY` (default `10s`): how long after a ban is made
|
||||||
|
`bans.json` is written, with every ban made in between.
|
||||||
|
- `SWWAF_STATE_COUNTER_INTERVAL` (default `15m`): how often every state file is
|
||||||
|
written.
|
||||||
|
|
||||||
Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are
|
Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are
|
||||||
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
|
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
|
||||||
@@ -193,12 +203,13 @@ ISO 3166-1 assigns today, and `xk` for Kosovo, in either case (`de` and `DE` are
|
|||||||
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
|
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
|
||||||
`su`, stops the start, and so does a code on both country lists. `off` switches
|
`su`, stops the start, and so does a code on both country lists. `off` switches
|
||||||
a timeout, a size limit or a rate limit off;
|
a timeout, a size limit or a rate limit off;
|
||||||
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` and the ban settings cannot be off.
|
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings and the state settings
|
||||||
|
cannot be off.
|
||||||
|
|
||||||
Several limits are fixed rather than settings. At most 20,000 clients are kept
|
Several limits are fixed rather than settings. At most 20,000 clients are kept,
|
||||||
for the rate limits, and an IPv6 client is counted by its /64. A new client
|
with their counters and history, and an IPv6 client is counted by its /64. A new
|
||||||
waits at most a second for its country, and at most 100,000 answers from GeoJS
|
client waits at most a second for its country, and at most 100,000 answers from
|
||||||
are kept, for 7 days each.
|
GeoJS are kept, for 7 days each.
|
||||||
|
|
||||||
## Request log
|
## Request log
|
||||||
|
|
||||||
@@ -246,6 +257,48 @@ which it answers `431`, headers slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`,
|
|||||||
whose connection it closes without an answer, and requests it cannot read at
|
whose connection it closes without an answer, and requests it cannot read at
|
||||||
all, which it answers itself, mostly with `400`.
|
all, which it answers itself, mostly with `400`.
|
||||||
|
|
||||||
|
## State files
|
||||||
|
|
||||||
|
`smallwebwaf` keeps its state in memory and a copy of it in three JSON files in
|
||||||
|
`SWWAF_STATE_DIR`, `/var/lib/smallwebwaf` by default, as "Persistent state" in
|
||||||
|
[`SPEC.md`](SPEC.md) describes. Each has a top-level `version`, 1, and lists its
|
||||||
|
entries by client address, with times in UTC.
|
||||||
|
|
||||||
|
- `bans.json`: every ban with its notes, indented to be read; a permanent ban's
|
||||||
|
`expires` is `null`.
|
||||||
|
- `clients.json`: each client's two buckets in the minute, the hour and the day,
|
||||||
|
and its history: when it was first and last seen, its country as last looked
|
||||||
|
up and when, its requests, how many were forwarded and how many refused, the
|
||||||
|
body bytes in each direction, its responses by status class and its offences
|
||||||
|
by kind. Each client is on a line of its own, so `grep` shows everything about
|
||||||
|
one.
|
||||||
|
- `lookups.json`: GeoJS's answers, one to a line, with when GeoJS gave each and
|
||||||
|
when it was last used.
|
||||||
|
|
||||||
|
`bans.json` is written `SWWAF_STATE_WRITE_DELAY` after a ban is made, with every
|
||||||
|
ban made in between, and every file every `SWWAF_STATE_COUNTER_INTERVAL` and
|
||||||
|
when `smallwebwaf` stops. Each write goes to a temporary file in the same
|
||||||
|
directory, which then replaces the file, so a crash leaves the old file or the
|
||||||
|
new one, whole. A write that fails is logged, and tried again at the next write.
|
||||||
|
A hard kill loses what changed since the last write.
|
||||||
|
|
||||||
|
At start the files are read back: each client keeps its counts, so a restart
|
||||||
|
gives it no fresh allowance, and each ban keeps refusing every client in its
|
||||||
|
netblock until it ends, even after `SWWAF_BAN_SCOPE_V4_PREFIX` has changed. A
|
||||||
|
netblock whose address has bits past its length, such as `203.0.113.9/24`, is
|
||||||
|
read as the netblock it is in, `203.0.113.0/24`. Buckets and answers whose time
|
||||||
|
has passed are dropped. A missing file is empty state, as on a first start. A
|
||||||
|
file that does not parse, or has another `version`, stops the start with a
|
||||||
|
message naming the file, and the line and column where Go's JSON decoder gives
|
||||||
|
them; so does a state directory `smallwebwaf` cannot write. So does an entry
|
||||||
|
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
|
||||||
|
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
|
||||||
|
`answered`. An edit made while `smallwebwaf` runs is overwritten by its next
|
||||||
|
write: taking it in comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
|
||||||
|
The AS number and AS name come with their lookup.
|
||||||
|
|
||||||
## Why
|
## Why
|
||||||
|
|
||||||
Small self-hosted sites now receive a great deal of traffic nobody asked for:
|
Small self-hosted sites now receive a great deal of traffic nobody asked for:
|
||||||
@@ -347,9 +400,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
|
||||||
come in milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md));
|
for the bans, the clients and the GeoJS answers are built (see "State files"
|
||||||
until then the rate counters, the bans and the GeoJS answers are kept in
|
above); the others come with their features, and taking in an edit while
|
||||||
memory only, and a restart loses them.
|
running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
|
||||||
- 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
|
||||||
@@ -445,8 +498,9 @@ main "$@"
|
|||||||
the health check on `127.0.0.1`.
|
the health check on `127.0.0.1`.
|
||||||
- `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf`. Mount a volume
|
- `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf`. Mount a volume
|
||||||
there to keep bans and client history when a deploy replaces the container;
|
there to keep bans and client history when a deploy replaces the container;
|
||||||
without one, it still starts. The state files come in milestone 3 or later;
|
without one, it still starts. At each start the `run` script of `smallwebwaf`
|
||||||
until then it writes nothing to disk and needs no volume.
|
gives that directory and every file in it to the `smallwebwaf` user, so a host
|
||||||
|
directory mounted there needs no change of owner.
|
||||||
- `docker stop` has runit stop both processes. `smallwebwaf` then stops taking
|
- `docker stop` has runit stop both processes. `smallwebwaf` then stops taking
|
||||||
requests and gives those in progress five seconds to finish.
|
requests and gives those in progress five seconds to finish.
|
||||||
|
|
||||||
@@ -484,11 +538,10 @@ client, for the request log, the metrics and the ban notes, and for the country
|
|||||||
lists and biased limits when you set them. It works with no setup: by default it
|
lists and biased limits when you set them. It works with no setup: by default it
|
||||||
asks the free GeoJS web service, which needs no account and no file. This means
|
asks the free GeoJS web service, which needs no account and no file. This means
|
||||||
that, by default, the address of every new visitor is sent to GeoJS. Each answer
|
that, by default, the address of every new visitor is sent to GeoJS. Each answer
|
||||||
is kept in memory for seven days, and many addresses are asked about in one
|
is kept for seven days, in memory and in `lookups.json`, so that it survives a
|
||||||
request; writing the answers to disk, so that they survive a restart, comes in
|
restart, and many addresses are asked about in one request. GeoJS publishes no
|
||||||
milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md)). GeoJS
|
rate limit but may block a caller it thinks asks too much; while it is not
|
||||||
publishes no rate limit but may block a caller it thinks asks too much; while it
|
answering, new visitors count as coming from an unknown country, which
|
||||||
is not answering, new visitors count as coming from an unknown country, which
|
|
||||||
`SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses.
|
`SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses.
|
||||||
|
|
||||||
To keep your visitors' addresses on your own host, set
|
To keep your visitors' addresses on your own host, set
|
||||||
@@ -519,9 +572,10 @@ addresses are never sent to GeoJS.
|
|||||||
## How the code is laid out
|
## How the code is laid out
|
||||||
|
|
||||||
- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`.
|
- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`.
|
||||||
- `internal/smallwebwaf`: the process: it reads the settings, listens, serves
|
- `internal/smallwebwaf`: the process: it reads the settings and the state
|
||||||
requests until `SIGTERM` or `SIGINT`, and stops. Run as
|
files, listens, serves requests until `SIGTERM` or `SIGINT`, and stops,
|
||||||
`smallwebwaf healthcheck`, it is the image's health check instead.
|
writing the state files. Run as `smallwebwaf healthcheck`, it is the image's
|
||||||
|
health check instead.
|
||||||
- `internal/config`: reads the settings, the one place they are read.
|
- `internal/config`: reads the settings, the one place they are read.
|
||||||
- `internal/proxy`: what happens to each request: it works out the client, runs
|
- `internal/proxy`: what happens to each request: it works out the client, runs
|
||||||
the checks, passes the request to the app and the answer back with the
|
the checks, passes the request to the app and the answer back with the
|
||||||
@@ -534,8 +588,10 @@ addresses are never sent to GeoJS.
|
|||||||
long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held.
|
long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held.
|
||||||
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the
|
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the
|
||||||
answers.
|
answers.
|
||||||
- `internal/ratelimit`: counts each client's requests and tells when one takes
|
- `internal/ratelimit`: the table of clients: counts each client's requests,
|
||||||
it over a rate limit.
|
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
|
||||||
|
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
|
||||||
@@ -576,20 +632,24 @@ so that they run in minimal containers.
|
|||||||
- `script/install-precommit`: installs that hook; `make hooks` runs it.
|
- `script/install-precommit`: installs that hook; `make hooks` runs it.
|
||||||
- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for
|
- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for
|
||||||
working on the code by hand; `make build` runs it.
|
working on the code by hand; `make build` runs it.
|
||||||
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it;
|
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it, with
|
||||||
`make run` runs it.
|
its state files in `bin/state` unless `SWWAF_STATE_DIR` is set; `make run`
|
||||||
|
runs it.
|
||||||
- `script/example-app`: builds the image and, on it, the example app in
|
- `script/example-app`: builds the image and, on it, the example app in
|
||||||
`deploy/example-app`, runs it, and checks that the health check passes, that a
|
`deploy/example-app`, runs it with a volume for the state files, and checks
|
||||||
request reaches the app through `smallwebwaf`, and that `sv stop` and
|
that the health check passes, that a request reaches the app through
|
||||||
`docker stop` stop it in order; then removes the container and both images. It
|
`smallwebwaf`, that a second request in a minute bans the client, that
|
||||||
needs network access, for nixpkgs' binary cache, and `script/check` does not
|
`sv stop` and `docker stop` stop it in order, and that a new container on the
|
||||||
run it; `make example-app` does.
|
same volume still refuses the banned client; then removes the containers, the
|
||||||
|
volume and both images. It needs network access, for nixpkgs' binary cache,
|
||||||
|
and `script/check` does not run it; `make example-app` does.
|
||||||
|
|
||||||
## TODO
|
## TODO
|
||||||
|
|
||||||
- The rest of milestone 3, after the bans that broken rate limits lead to and up
|
- The rest of milestone 3, from taking in an admin's edits to the state files
|
||||||
to the metrics endpoint, and the rest of the design, in the order of the build
|
(https://git.eeqj.de/sneak/smallwebwaf/issues/68) up to the metrics endpoint,
|
||||||
order in [`SPEC.md`](SPEC.md).
|
and the rest of the design, in the order of the build order in
|
||||||
|
[`SPEC.md`](SPEC.md).
|
||||||
|
|
||||||
## Documents
|
## Documents
|
||||||
|
|
||||||
|
|||||||
+142
-46
@@ -1,6 +1,7 @@
|
|||||||
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
||||||
// netblocks of clients that break a rate limit, with their notes, as the
|
// netblocks of clients that break a rate limit, with their notes, as the
|
||||||
// "Bans" section of SPEC.md describes. The bans are kept in memory only.
|
// "Bans" section of SPEC.md describes. The bans are kept in memory, and
|
||||||
|
// written to bans.json and read from it by the state package.
|
||||||
package bans
|
package bans
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -58,48 +59,65 @@ func (b Ban) ActiveAt(now time.Time) bool {
|
|||||||
return b.Permanent() || now.Before(b.Expires)
|
return b.Permanent() || now.Before(b.Expires)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Notes are what an admin needs to decide whether to lift a ban.
|
// Notes are what an admin needs to decide whether to lift a ban. The
|
||||||
|
// JSON names are those of bans.json.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
type Notes struct {
|
type Notes struct {
|
||||||
// Country is the client's country, when it was looked up.
|
// Country is the client's country, when it was looked up.
|
||||||
Country string
|
Country string `json:"country"`
|
||||||
// Limit, Window and Count are the limit that was broken, its window,
|
// Limit, Window and Count are the limit that was broken, its window,
|
||||||
// "minute", "hour" or "day", and the count reached: the client's
|
// "minute", "hour" or "day", and the count reached: the client's
|
||||||
// requests in the window, the one that broke the limit included.
|
// requests in the window, the one that broke the limit included.
|
||||||
// These are the requests that counted toward the ban, and the window
|
// These are the requests that counted toward the ban, and the window
|
||||||
// is the time over which they came.
|
// is the time over which they came.
|
||||||
Limit int64
|
Limit int64 `json:"limit"`
|
||||||
Window string
|
Window string `json:"window"`
|
||||||
Count float64
|
Count float64 `json:"count"`
|
||||||
// Request is the request that broke the limit.
|
// Request is the request that broke the limit.
|
||||||
Request Request
|
Request Request `json:"request"`
|
||||||
// Refused is how many requests the ban has refused so far.
|
// Requests is how many requests the netblock has sent since it was
|
||||||
Refused int64
|
// first seen, and Refused how many of them the ban has refused so
|
||||||
|
// far. Both go up with each request the ban refuses.
|
||||||
|
Requests int64 `json:"requests"`
|
||||||
|
Refused int64 `json:"refused"`
|
||||||
// EarlierBans is how many bans the netblock had before this one.
|
// EarlierBans is how many bans the netblock had before this one.
|
||||||
EarlierBans int
|
EarlierBans int `json:"earlier_bans"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Request is a request in a ban's notes. Each text is cut to 256 bytes.
|
// Request is a request in a ban's notes. Each text is cut to 256 bytes.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
type Request struct {
|
type Request struct {
|
||||||
Time time.Time
|
Time time.Time `json:"time"`
|
||||||
Method string
|
Method string `json:"method"`
|
||||||
Host string
|
Host string `json:"host"`
|
||||||
// Path is the path with its query string.
|
// Path is the path with its query string.
|
||||||
Path string
|
Path string `json:"path"`
|
||||||
// Status is what the client was sent, 0 if nothing was.
|
// Status is what the client was sent, 0 if nothing was.
|
||||||
Status int
|
Status int `json:"status"`
|
||||||
UserAgent string
|
UserAgent string `json:"user_agent"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Ledger holds the bans. It is safe for concurrent use.
|
// Ledger holds the bans. It is safe for concurrent use.
|
||||||
type Ledger struct {
|
type Ledger struct {
|
||||||
rules Rules
|
rules Rules
|
||||||
|
// changed receives a value when a ban is made, unless one is waiting
|
||||||
|
// already.
|
||||||
|
changed chan struct{}
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
// netblocks holds each banned netblock's bans, oldest first. Each
|
// netblocks holds each banned netblock's bans, oldest first. Check
|
||||||
// request from a netblock makes it the most recently seen.
|
// makes each netblock it finds the most recently seen.
|
||||||
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
|
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
|
||||||
// held is how many bans netblocks holds, at most rules.MaxBans.
|
// held is how many bans netblocks holds, at most rules.MaxBans.
|
||||||
held int
|
held int
|
||||||
|
// v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6
|
||||||
|
// netblocks that have been banned. Check looks for a ban at each of
|
||||||
|
// them, so that a ban read from bans.json refuses every client in its
|
||||||
|
// netblock even when it was made with another SWWAF_BAN_SCOPE_V4_PREFIX,
|
||||||
|
// or another length of an IPv6 client's netblock.
|
||||||
|
v4Lengths, v6Lengths []int
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns a Ledger with no ban yet.
|
// New returns a Ledger with no ban yet.
|
||||||
@@ -111,31 +129,49 @@ func New(rules Rules) *Ledger {
|
|||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Ledger{rules: rules, netblocks: netblocks}
|
return &Ledger{
|
||||||
|
rules: rules,
|
||||||
|
changed: make(chan struct{}, 1),
|
||||||
|
netblocks: netblocks,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check is called for each request from netblock, at now. It reports
|
// Changed receives a value after a ban is made, so that bans.json can be
|
||||||
// whether a ban on netblock is active, and returns that ban, with the
|
// written. Several bans made before it is read leave one value.
|
||||||
// request counted among those it refused.
|
func (l *Ledger) Changed() <-chan struct{} {
|
||||||
func (l *Ledger) Check(netblock netip.Prefix, now time.Time) (Ban, bool) {
|
return l.changed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check is called for each request from client, at now. It reports
|
||||||
|
// whether a ban on a netblock client is in is active, and returns that
|
||||||
|
// ban, with the request counted among those it refused.
|
||||||
|
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
bans, found := l.netblocks.Get(netblock)
|
lengths := l.v6Lengths
|
||||||
if !found {
|
if client.Is4() {
|
||||||
return Ban{}, false
|
lengths = l.v4Lengths
|
||||||
}
|
}
|
||||||
|
|
||||||
// A ban is made only once the one before has ended, so only the last
|
for _, length := range lengths {
|
||||||
// can be active.
|
bans, found := l.netblocks.Get(netip.PrefixFrom(client, length).Masked())
|
||||||
last := &(*bans)[len(*bans)-1]
|
if !found {
|
||||||
if !last.ActiveAt(now) {
|
continue
|
||||||
return Ban{}, false
|
}
|
||||||
|
|
||||||
|
// A ban is made only once the one before has ended, so only the
|
||||||
|
// last can be active.
|
||||||
|
last := &(*bans)[len(*bans)-1]
|
||||||
|
if last.ActiveAt(now) {
|
||||||
|
last.Notes.Requests++
|
||||||
|
last.Notes.Refused++
|
||||||
|
|
||||||
|
return *last, true
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
last.Notes.Refused++
|
return Ban{}, false
|
||||||
|
|
||||||
return *last, true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// BanForLimit bans netblock at now for a broken limit, with notes, and
|
// BanForLimit bans netblock at now for a broken limit, with notes, and
|
||||||
@@ -169,21 +205,13 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
|
|||||||
Expires: l.expiry(last, now),
|
Expires: l.expiry(last, now),
|
||||||
Notes: notes,
|
Notes: notes,
|
||||||
}
|
}
|
||||||
|
l.add(ban)
|
||||||
|
|
||||||
if l.held == l.rules.MaxBans {
|
select {
|
||||||
l.dropOne()
|
case l.changed <- struct{}{}:
|
||||||
|
default: // a value is waiting already
|
||||||
}
|
}
|
||||||
|
|
||||||
// dropOne can have dropped netblock's last ban, and netblock with it.
|
|
||||||
bans, found = l.netblocks.Peek(netblock)
|
|
||||||
if !found {
|
|
||||||
bans = &[]Ban{}
|
|
||||||
l.netblocks.Add(netblock, bans)
|
|
||||||
}
|
|
||||||
|
|
||||||
*bans = append(*bans, ban)
|
|
||||||
l.held++
|
|
||||||
|
|
||||||
return ban
|
return ban
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -201,6 +229,74 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
|
|||||||
return slices.Clone(*bans)
|
return slices.Clone(*bans)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Snapshot returns every ban held, sorted by netblock, and each
|
||||||
|
// netblock's bans oldest first, as bans.json lists them.
|
||||||
|
func (l *Ledger) Snapshot() []Ban {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
held := make([]Ban, 0, l.held)
|
||||||
|
for _, bans := range l.netblocks.Values() {
|
||||||
|
held = append(held, *bans...)
|
||||||
|
}
|
||||||
|
|
||||||
|
slices.SortStableFunc(held, func(a, b Ban) int {
|
||||||
|
return a.Netblock.Compare(b.Netblock)
|
||||||
|
})
|
||||||
|
|
||||||
|
return held
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load puts bans read from bans.json into a ledger that holds none yet,
|
||||||
|
// in the order they started, so that a netblock whose last ban started
|
||||||
|
// latest counts as the most recently seen. Each netblock is masked to its
|
||||||
|
// length, so that 203.0.113.9/24 is 203.0.113.0/24, and each text in the
|
||||||
|
// notes is cut to 256 bytes. Past MaxBans the earliest bans are dropped,
|
||||||
|
// as when they are made.
|
||||||
|
func (l *Ledger) Load(bans []Ban) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
bans = slices.Clone(bans)
|
||||||
|
slices.SortStableFunc(bans, func(a, b Ban) int {
|
||||||
|
return a.Start.Compare(b.Start)
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, ban := range bans {
|
||||||
|
ban.Netblock = ban.Netblock.Masked()
|
||||||
|
ban.Notes.Request = ban.Notes.Request.cut()
|
||||||
|
l.add(ban)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// add adds ban to its netblock's bans, after the last, and makes its
|
||||||
|
// netblock the most recently seen. With MaxBans held, it drops one first.
|
||||||
|
func (l *Ledger) add(ban Ban) {
|
||||||
|
if l.held == l.rules.MaxBans {
|
||||||
|
l.dropOne()
|
||||||
|
}
|
||||||
|
|
||||||
|
// dropOne can have dropped the netblock's last ban, and the netblock
|
||||||
|
// with it.
|
||||||
|
bans, found := l.netblocks.Get(ban.Netblock)
|
||||||
|
if !found {
|
||||||
|
bans = &[]Ban{}
|
||||||
|
l.netblocks.Add(ban.Netblock, bans)
|
||||||
|
}
|
||||||
|
|
||||||
|
*bans = append(*bans, ban)
|
||||||
|
l.held++
|
||||||
|
|
||||||
|
lengths := &l.v6Lengths
|
||||||
|
if ban.Netblock.Addr().Is4() {
|
||||||
|
lengths = &l.v4Lengths
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Contains(*lengths, ban.Netblock.Bits()) {
|
||||||
|
*lengths = append(*lengths, ban.Netblock.Bits())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 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 last ban, which has ended,
|
||||||
// or nil when it has none.
|
// or nil when it has none.
|
||||||
|
|||||||
+12
-10
@@ -39,7 +39,7 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
|
|||||||
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
|
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned := ledger.Check(netblock, now.Add(100*365*day))
|
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day))
|
||||||
if !banned {
|
if !banned {
|
||||||
t.Error("a permanent ban ended")
|
t.Error("a permanent ban ended")
|
||||||
}
|
}
|
||||||
@@ -135,28 +135,30 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
|
|||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
ledger := bans.New(defaultRules())
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
||||||
|
|
||||||
for range 3 {
|
for range 3 {
|
||||||
got, banned := ledger.Check(netblock, ban.Expires.Add(-time.Nanosecond))
|
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
||||||
if !banned || got.Start != ban.Start {
|
if !banned || got.Start != ban.Start {
|
||||||
t.Fatalf("check during the ban gives %+v and %t", got, banned)
|
t.Fatalf("check during the ban gives %+v and %t", got, banned)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned := ledger.Check(netip.MustParsePrefix("203.0.113.10/32"), midnight())
|
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("another netblock is banned")
|
t.Error("another netblock is banned")
|
||||||
}
|
}
|
||||||
|
|
||||||
_, banned = ledger.Check(netblock, ban.Expires)
|
_, banned = ledger.Check(netblock.Addr(), ban.Expires)
|
||||||
if banned {
|
if banned {
|
||||||
t.Error("the ban did not end")
|
t.Error("the ban did not end")
|
||||||
}
|
}
|
||||||
|
|
||||||
refused := ledger.Bans(netblock)[0].Notes.Refused
|
// The netblock's requests went from 5 to 8 with the three refused.
|
||||||
if refused != 3 {
|
notes := ledger.Bans(netblock)[0].Notes
|
||||||
t.Errorf("the notes count %d refused requests, want 3", refused)
|
if notes.Refused != 3 || notes.Requests != 8 {
|
||||||
|
t.Errorf("the notes count %d refused requests of %d, want 3 of 8",
|
||||||
|
notes.Refused, notes.Requests)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -178,7 +180,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
|
|||||||
|
|
||||||
// A request from a makes b the netblock seen longest ago, and its ban
|
// A request from a makes b the netblock seen longest ago, and its ban
|
||||||
// goes to make room for d's.
|
// goes to make room for d's.
|
||||||
ledger.Check(a, now)
|
ledger.Check(a.Addr(), now)
|
||||||
ledger.BanForLimit(d, now, bans.Notes{})
|
ledger.BanForLimit(d, now, bans.Notes{})
|
||||||
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
|
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
|
||||||
|
|
||||||
@@ -188,7 +190,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
|
|||||||
|
|
||||||
// With d seen since, a is seen longest ago, and its earlier ban goes
|
// With d seen since, a is seen longest ago, and its earlier ban goes
|
||||||
// first.
|
// first.
|
||||||
ledger.Check(d, first.Expires)
|
ledger.Check(d.Addr(), first.Expires)
|
||||||
ledger.BanForLimit(b, first.Expires, bans.Notes{})
|
ledger.BanForLimit(b, first.Expires, bans.Notes{})
|
||||||
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
|
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,193 @@
|
|||||||
|
package bans_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestChangedAfterABanIsMade(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
wantChanged(t, ledger, false)
|
||||||
|
|
||||||
|
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
|
// A limit broken during the ban makes no other, and a refusal changes
|
||||||
|
// only the counts in the notes, which wait for the interval's write.
|
||||||
|
ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
|
||||||
|
ledger.Check(netblock.Addr(), midnight().Add(time.Minute))
|
||||||
|
wantChanged(t, ledger, false)
|
||||||
|
|
||||||
|
// Two bans before the value is read leave one.
|
||||||
|
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{})
|
||||||
|
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.11/32"), midnight(), bans.Notes{})
|
||||||
|
wantChanged(t, ledger, true)
|
||||||
|
wantChanged(t, ledger, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
v6 := netip.MustParsePrefix("2001:db8::/64")
|
||||||
|
high := netip.MustParsePrefix("203.0.113.10/32")
|
||||||
|
low := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
first := ledger.BanForLimit(v6, midnight(), bans.Notes{})
|
||||||
|
ledger.BanForLimit(high, midnight(), bans.Notes{})
|
||||||
|
ledger.BanForLimit(low, midnight(), bans.Notes{})
|
||||||
|
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
|
||||||
|
|
||||||
|
snapshot := ledger.Snapshot()
|
||||||
|
|
||||||
|
got := make([]string, 0, len(snapshot))
|
||||||
|
for _, ban := range snapshot {
|
||||||
|
got = append(got, ban.Netblock.String()+" "+ban.Start.Format(time.Kitchen))
|
||||||
|
}
|
||||||
|
|
||||||
|
want := []string{
|
||||||
|
"203.0.113.9/32 12:00AM", "203.0.113.10/32 12:00AM",
|
||||||
|
"2001:db8::/64 12:00AM", "2001:db8::/64 1:00AM",
|
||||||
|
}
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("snapshot %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadedBansCarryOn(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
before := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
|
||||||
|
|
||||||
|
// Loaded into a new ledger, as across a restart, the ban still refuses
|
||||||
|
// while it lasts, and once it has ended a broken limit bans for three
|
||||||
|
// times as long, with the loaded ban counted among the earlier ones.
|
||||||
|
after := bans.New(defaultRules())
|
||||||
|
after.Load(before.Snapshot())
|
||||||
|
|
||||||
|
_, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
|
||||||
|
if !banned {
|
||||||
|
t.Error("the loaded ban does not refuse")
|
||||||
|
}
|
||||||
|
|
||||||
|
again := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
|
||||||
|
if again.Expires.Sub(again.Start) != 3*time.Hour || again.Notes.EarlierBans != 1 {
|
||||||
|
t.Errorf("the next ban lasts %s with %d earlier bans, want 3h and 1",
|
||||||
|
again.Expires.Sub(again.Start), again.Notes.EarlierBans)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Two entries as an admin might write them, with addresses not masked
|
||||||
|
// to their lengths, the IPv6 one shorter than the /64 an IPv6 client's
|
||||||
|
// ban covers, beside a ban the ledger makes on one IPv4 address.
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{
|
||||||
|
{Netblock: netip.MustParsePrefix("203.0.113.9/24"), Start: midnight()},
|
||||||
|
{Netblock: netip.MustParsePrefix("2001:db8::1/48"), Start: midnight()},
|
||||||
|
})
|
||||||
|
ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), bans.Notes{})
|
||||||
|
|
||||||
|
for client, want := range map[string]bool{
|
||||||
|
"203.0.113.0": true,
|
||||||
|
"203.0.113.200": true,
|
||||||
|
"203.0.114.1": false,
|
||||||
|
"2001:db8:0:5::1": true,
|
||||||
|
"2001:db8:1::1": false,
|
||||||
|
"198.51.100.7": true,
|
||||||
|
"198.51.100.8": false,
|
||||||
|
} {
|
||||||
|
_, banned := ledger.Check(netip.MustParseAddr(client), midnight())
|
||||||
|
if banned != want {
|
||||||
|
t.Errorf("%s is refused: %t, want %t", client, banned, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The loaded netblocks are written back masked.
|
||||||
|
snapshot := ledger.Snapshot()
|
||||||
|
|
||||||
|
got := make([]string, 0, len(snapshot))
|
||||||
|
for _, ban := range snapshot {
|
||||||
|
got = append(got, ban.Netblock.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
want := []string{"198.51.100.7/32", "203.0.113.0/24", "2001:db8::/48"}
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("the ledger holds bans on %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// bans.json lists the bans by netblock, not in the order they began.
|
||||||
|
later := bans.Ban{Netblock: netip.MustParsePrefix("203.0.113.1/32"), Start: midnight()}
|
||||||
|
earlier := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
||||||
|
Start: midnight().Add(-time.Hour),
|
||||||
|
}
|
||||||
|
|
||||||
|
rules := defaultRules()
|
||||||
|
rules.MaxBans = 1
|
||||||
|
ledger := bans.New(rules)
|
||||||
|
ledger.Load([]bans.Ban{later, earlier})
|
||||||
|
|
||||||
|
held := ledger.Snapshot()
|
||||||
|
if len(held) != 1 || held[0] != later {
|
||||||
|
t.Errorf("the ledger holds %+v, want only the ban that began later", held)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
long := strings.Repeat("a", 300)
|
||||||
|
ban := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.9/32"),
|
||||||
|
Start: midnight(),
|
||||||
|
Notes: bans.Notes{Request: bans.Request{
|
||||||
|
Method: long, Host: long, Path: long, UserAgent: long,
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{ban})
|
||||||
|
|
||||||
|
cut := long[:256]
|
||||||
|
want := bans.Request{Method: cut, Host: cut, Path: cut, UserAgent: cut}
|
||||||
|
|
||||||
|
got := ledger.Snapshot()[0].Notes.Request
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("the notes keep %+v, want each text cut to 256 bytes", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantChanged checks whether the ledger's Changed has a value to read.
|
||||||
|
func wantChanged(t *testing.T, ledger *bans.Ledger, want bool) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := false
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ledger.Changed():
|
||||||
|
got = true
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("Changed has a value: %t, want %t", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"path/filepath"
|
||||||
"slices"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -94,6 +95,14 @@ type Config struct {
|
|||||||
// BanScopeV4Prefix is the length of the netblock around an IPv4
|
// BanScopeV4Prefix is the length of the netblock around an IPv4
|
||||||
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
|
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
|
||||||
BanScopeV4Prefix int
|
BanScopeV4Prefix int
|
||||||
|
// StateDir is the directory of the state files, an absolute path
|
||||||
|
// (SWWAF_STATE_DIR). bans.json is written StateWriteDelay after a ban
|
||||||
|
// is made (SWWAF_STATE_WRITE_DELAY), and every state file every
|
||||||
|
// StateCounterInterval (SWWAF_STATE_COUNTER_INTERVAL). Neither can be
|
||||||
|
// off.
|
||||||
|
StateDir string
|
||||||
|
StateWriteDelay time.Duration
|
||||||
|
StateCounterInterval time.Duration
|
||||||
|
|
||||||
// settings are the values read, as given or by default, for the
|
// settings are the values read, as given or by default, for the
|
||||||
// log line at start.
|
// log line at start.
|
||||||
@@ -139,6 +148,8 @@ var (
|
|||||||
errNotBanResponse = errors.New("is not 403, 429 or close")
|
errNotBanResponse = errors.New("is not 403, 429 or close")
|
||||||
errNotV4Prefix = errors.New(
|
errNotV4Prefix = errors.New(
|
||||||
"is not the length of an IPv4 netblock, from 0 to 32, such as 24")
|
"is not the length of an IPv4 netblock, from 0 to 32, such as 24")
|
||||||
|
errNotAbsolutePath = errors.New(
|
||||||
|
"is not an absolute path, such as /var/lib/smallwebwaf")
|
||||||
)
|
)
|
||||||
|
|
||||||
// FromEnvironment reads the settings with lookupEnv, normally
|
// FromEnvironment reads the settings with lookupEnv, normally
|
||||||
@@ -174,6 +185,9 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
|||||||
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
|
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
|
||||||
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
|
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
|
||||||
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
|
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
|
||||||
|
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
|
||||||
|
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
|
||||||
|
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, country := range cfg.ExclusivelyAllowedCountries {
|
for _, country := range cfg.ExclusivelyAllowedCountries {
|
||||||
@@ -329,6 +343,16 @@ func (e *environment) v4Prefix(name, defaultValue string) int {
|
|||||||
return length
|
return length
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// absolutePath reads a setting that is an absolute path.
|
||||||
|
func (e *environment) absolutePath(name, defaultValue string) string {
|
||||||
|
path := e.value(name, defaultValue)
|
||||||
|
if !filepath.IsAbs(path) {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", path, errNotAbsolutePath))
|
||||||
|
}
|
||||||
|
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
||||||
// whole number of days such as 7d, or off.
|
// whole number of days such as 7d, or off.
|
||||||
func parseDuration(value string) (time.Duration, error) {
|
func parseDuration(value string) (time.Duration, error) {
|
||||||
|
|||||||
@@ -41,6 +41,9 @@ 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"
|
||||||
|
stateDir = "SWWAF_STATE_DIR"
|
||||||
|
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
|
||||||
|
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
||||||
)
|
)
|
||||||
|
|
||||||
// off switches a timeout, a size limit or a rate limit off.
|
// off switches a timeout, a size limit or a rate limit off.
|
||||||
@@ -92,6 +95,9 @@ func TestDefaults(t *testing.T) {
|
|||||||
MaxBanDuration: 7 * 24 * time.Hour,
|
MaxBanDuration: 7 * 24 * time.Hour,
|
||||||
MaxBans: 5000,
|
MaxBans: 5000,
|
||||||
BanScopeV4Prefix: 32,
|
BanScopeV4Prefix: 32,
|
||||||
|
StateDir: "/var/lib/smallwebwaf",
|
||||||
|
StateWriteDelay: 10 * time.Second,
|
||||||
|
StateCounterInterval: 15 * time.Minute,
|
||||||
})
|
})
|
||||||
|
|
||||||
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
||||||
@@ -136,6 +142,9 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
maxBanDuration: "30d",
|
maxBanDuration: "30d",
|
||||||
maxBans: "100",
|
maxBans: "100",
|
||||||
banScopeV4Prefix: "24",
|
banScopeV4Prefix: "24",
|
||||||
|
stateDir: "/srv/waf-state",
|
||||||
|
stateWriteDelay: "500ms",
|
||||||
|
stateCounterInterval: "1h",
|
||||||
})
|
})
|
||||||
|
|
||||||
wantSettings(t, cfg, config.Config{
|
wantSettings(t, cfg, config.Config{
|
||||||
@@ -157,6 +166,9 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
MaxBanDuration: 30 * 24 * time.Hour,
|
MaxBanDuration: 30 * 24 * time.Hour,
|
||||||
MaxBans: 100,
|
MaxBans: 100,
|
||||||
BanScopeV4Prefix: 24,
|
BanScopeV4Prefix: 24,
|
||||||
|
StateDir: "/srv/waf-state",
|
||||||
|
StateWriteDelay: 500 * time.Millisecond,
|
||||||
|
StateCounterInterval: time.Hour,
|
||||||
})
|
})
|
||||||
|
|
||||||
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
|
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
|
||||||
@@ -329,6 +341,9 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{maxBanDuration, off}, {maxBanDuration, "1w"},
|
{maxBanDuration, off}, {maxBanDuration, "1w"},
|
||||||
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
|
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
|
||||||
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
|
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
|
||||||
|
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
|
||||||
|
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
|
||||||
|
{stateCounterInterval, off}, {stateCounterInterval, "15"},
|
||||||
} {
|
} {
|
||||||
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
@@ -389,6 +404,9 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
maxBanDuration: "7d",
|
maxBanDuration: "7d",
|
||||||
maxBans: "5000",
|
maxBans: "5000",
|
||||||
banScopeV4Prefix: "32",
|
banScopeV4Prefix: "32",
|
||||||
|
stateDir: "/var/lib/smallwebwaf",
|
||||||
|
stateWriteDelay: "10s",
|
||||||
|
stateCounterInterval: "15m",
|
||||||
}
|
}
|
||||||
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)
|
||||||
@@ -417,7 +435,7 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
|||||||
wantBanSettings(t, got, want)
|
wantBanSettings(t, got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantBanSettings checks the settings for bans.
|
// wantBanSettings checks the settings for bans and the state files.
|
||||||
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
|
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
@@ -429,6 +447,12 @@ func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
|
|||||||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
|
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
|
||||||
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
|
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if got.StateDir != want.StateDir ||
|
||||||
|
got.StateWriteDelay != want.StateWriteDelay ||
|
||||||
|
got.StateCounterInterval != want.StateCounterInterval {
|
||||||
|
t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantNetblocks checks a list of netblocks.
|
// wantNetblocks checks a list of netblocks.
|
||||||
|
|||||||
+65
-13
@@ -1,6 +1,7 @@
|
|||||||
// Package lookup looks up each client's country through the GeoJS web
|
// Package lookup looks up each client's country through the GeoJS web
|
||||||
// service, and keeps the answers in memory, for at most 100,000 clients
|
// service, and keeps the answers in memory, for at most 100,000 clients
|
||||||
// and for 7 days each.
|
// and for 7 days each. The answers are written to lookups.json and read
|
||||||
|
// from it by the state package.
|
||||||
package lookup
|
package lookup
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -12,6 +13,7 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -76,7 +78,7 @@ type GeoJS struct {
|
|||||||
httpClient *http.Client
|
httpClient *http.Client
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
answers *simplelru.LRU[netip.Prefix, answer]
|
answers *simplelru.LRU[netip.Prefix, *Answer]
|
||||||
// waiting are the clients without an answer: those to ask GeoJS about,
|
// waiting are the clients without an answer: those to ask GeoJS about,
|
||||||
// and those it is being asked about.
|
// and those it is being asked about.
|
||||||
waiting map[netip.Prefix]*wait
|
waiting map[netip.Prefix]*wait
|
||||||
@@ -88,11 +90,14 @@ type GeoJS struct {
|
|||||||
retryAt time.Time
|
retryAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
// answer is what GeoJS said about a client: its country, "" when GeoJS
|
// Answer is what GeoJS said about a client, as lookups.json holds it: its
|
||||||
// cannot place it, and when GeoJS said so.
|
// country, "" when GeoJS cannot place it, when GeoJS said so, and when
|
||||||
type answer struct {
|
// the answer was last used.
|
||||||
country string
|
type Answer struct {
|
||||||
received time.Time
|
Client netip.Prefix `json:"client"`
|
||||||
|
Country string `json:"country"`
|
||||||
|
Answered time.Time `json:"answered"`
|
||||||
|
Used time.Time `json:"used"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// wait is a client waiting for its answer.
|
// wait is a client waiting for its answer.
|
||||||
@@ -107,7 +112,7 @@ type wait struct {
|
|||||||
|
|
||||||
// New returns a GeoJS with no answer kept yet.
|
// New returns a GeoJS with no answer kept yet.
|
||||||
func New(params Params) *GeoJS {
|
func New(params Params) *GeoJS {
|
||||||
answers, err := simplelru.NewLRU[netip.Prefix, answer](maxAnswers, nil)
|
answers, err := simplelru.NewLRU[netip.Prefix, *Answer](maxAnswers, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
@@ -164,6 +169,47 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
|||||||
return country
|
return country
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Snapshot returns every answer kept, sorted by client, as lookups.json
|
||||||
|
// lists them.
|
||||||
|
func (g *GeoJS) Snapshot() []Answer {
|
||||||
|
g.mu.Lock()
|
||||||
|
|
||||||
|
answers := make([]Answer, 0, g.answers.Len())
|
||||||
|
for _, kept := range g.answers.Values() {
|
||||||
|
answers = append(answers, *kept)
|
||||||
|
}
|
||||||
|
|
||||||
|
g.mu.Unlock()
|
||||||
|
|
||||||
|
slices.SortFunc(answers, func(a, b Answer) int {
|
||||||
|
return a.Client.Compare(b.Client)
|
||||||
|
})
|
||||||
|
|
||||||
|
return answers
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load keeps answers read from lookups.json, in a GeoJS that keeps none
|
||||||
|
// yet, in the order they were last used, so that the one used longest
|
||||||
|
// ago is dropped first. Answers GeoJS gave keepFor ago or more are
|
||||||
|
// dropped.
|
||||||
|
func (g *GeoJS) Load(answers []Answer) {
|
||||||
|
g.mu.Lock()
|
||||||
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
|
answers = slices.Clone(answers)
|
||||||
|
slices.SortStableFunc(answers, func(a, b Answer) int {
|
||||||
|
return a.Used.Compare(b.Used)
|
||||||
|
})
|
||||||
|
|
||||||
|
now := g.now()
|
||||||
|
|
||||||
|
for _, answer := range answers {
|
||||||
|
if now.Sub(answer.Answered) < keepFor {
|
||||||
|
g.answers.Add(answer.Client, &answer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// answerOrWait returns client's kept answer if it has one. Otherwise it
|
// answerOrWait returns client's kept answer if it has one. Otherwise it
|
||||||
// puts the client among those waiting if there is room, has GeoJS asked
|
// puts the client among those waiting if there is room, has GeoJS asked
|
||||||
// about them if it can be, and returns what to wait on for the answer, or
|
// about them if it can be, and returns what to wait on for the answer, or
|
||||||
@@ -203,15 +249,19 @@ func (g *GeoJS) answerOrWait(
|
|||||||
return "", w.asked
|
return "", w.asked
|
||||||
}
|
}
|
||||||
|
|
||||||
// kept returns client's answer, if one was received less than keepFor
|
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
|
||||||
// ago.
|
// and notes that it was used.
|
||||||
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
|
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
|
||||||
|
now := g.now()
|
||||||
|
|
||||||
kept, found := g.answers.Get(client)
|
kept, found := g.answers.Get(client)
|
||||||
if !found || g.now().Sub(kept.received) >= keepFor {
|
if !found || now.Sub(kept.Answered) >= keepFor {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
return kept.country, true
|
kept.Used = now
|
||||||
|
|
||||||
|
return kept.Country, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// ask starts asking GeoJS about the waiting clients, unless a request to
|
// ask starts asking GeoJS about the waiting clients, unless a request to
|
||||||
@@ -293,7 +343,9 @@ func (g *GeoJS) keep(
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
g.answers.Add(client, answer{country: country, received: now})
|
g.answers.Add(client, &Answer{
|
||||||
|
Client: client, Country: country, Answered: now, Used: now,
|
||||||
|
})
|
||||||
close(g.waiting[client].asked)
|
close(g.waiting[client].asked)
|
||||||
delete(g.waiting, client)
|
delete(g.waiting, client)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package lookup_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, clock, g := start(t)
|
||||||
|
placed := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
notPlaced := netip.MustParsePrefix(unplaced + "/32")
|
||||||
|
asked := clock.Now()
|
||||||
|
|
||||||
|
wantCountry(t, g, placed, germany)
|
||||||
|
wantCountry(t, g, notPlaced, "")
|
||||||
|
|
||||||
|
clock.advance(time.Hour)
|
||||||
|
wantCountry(t, g, placed, germany)
|
||||||
|
|
||||||
|
want := []lookup.Answer{
|
||||||
|
{Client: notPlaced, Country: "", Answered: asked, Used: asked},
|
||||||
|
{Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
|
||||||
|
}
|
||||||
|
if got := g.Snapshot(); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadedAnswersAreKeptFor7DaysFromWhenGeoJSGaveThem(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojs, clock, g := start(t)
|
||||||
|
now := clock.Now()
|
||||||
|
kept := lookup.Answer{
|
||||||
|
Client: netip.MustParsePrefix("203.0.113.9/32"),
|
||||||
|
Country: "FR",
|
||||||
|
Answered: now.Add(-week + time.Second),
|
||||||
|
Used: now.Add(-time.Hour),
|
||||||
|
}
|
||||||
|
stale := lookup.Answer{
|
||||||
|
Client: netip.MustParsePrefix("203.0.113.10/32"),
|
||||||
|
Country: "FR",
|
||||||
|
Answered: now.Add(-week),
|
||||||
|
Used: now.Add(-time.Hour),
|
||||||
|
}
|
||||||
|
|
||||||
|
g.Load([]lookup.Answer{kept, stale})
|
||||||
|
|
||||||
|
if got := g.Snapshot(); !slices.Equal(got, []lookup.Answer{kept}) {
|
||||||
|
t.Errorf("kept %+v, want only the answer GeoJS gave less than 7 days ago", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantCountry(t, g, kept.Client, "FR")
|
||||||
|
wantRequests(t, geojs, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadDropsTheAnswerUsedLongestAgoFirst(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const maxAnswers = 100000
|
||||||
|
|
||||||
|
_, clock, g := start(t)
|
||||||
|
now := clock.Now()
|
||||||
|
|
||||||
|
// lookups.json lists the answers by client. Here each was last used a
|
||||||
|
// second before the one listed before it, so the last listed is the
|
||||||
|
// one used longest ago, and the one dropped.
|
||||||
|
answers := make([]lookup.Answer, maxAnswers+1)
|
||||||
|
addr := netip.MustParseAddr("10.0.0.0")
|
||||||
|
|
||||||
|
for i := range answers {
|
||||||
|
answers[i] = lookup.Answer{
|
||||||
|
Client: netip.PrefixFrom(addr, addr.BitLen()),
|
||||||
|
Country: germany,
|
||||||
|
Answered: now,
|
||||||
|
Used: now.Add(-time.Duration(i) * time.Second),
|
||||||
|
}
|
||||||
|
addr = addr.Next()
|
||||||
|
}
|
||||||
|
|
||||||
|
g.Load(answers)
|
||||||
|
|
||||||
|
got := g.Snapshot()
|
||||||
|
if len(got) != maxAnswers || got[0] != answers[0] ||
|
||||||
|
got[maxAnswers-1] != answers[maxAnswers-1] {
|
||||||
|
t.Errorf("%d answers kept, from %s to %s; want %d, from %s to %s",
|
||||||
|
len(got), got[0].Client, got[len(got)-1].Client, maxAnswers,
|
||||||
|
answers[0].Client, answers[maxAnswers-1].Client)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -14,10 +14,10 @@ func (rq *request) banResponse(action string) *refusal {
|
|||||||
return &refusal{status: rq.h.config.BanResponse, action: action}
|
return &refusal{status: rq.h.config.BanResponse, action: action}
|
||||||
}
|
}
|
||||||
|
|
||||||
// banned reports whether a ban on the client's netblock refuses the
|
// banned reports whether a ban on a netblock the client is in refuses
|
||||||
// request at now, and notes for the log line when that ban ends.
|
// the request at now, and notes for the log line when that ban ends.
|
||||||
func (rq *request) banned(now time.Time) bool {
|
func (rq *request) banned(now time.Time) bool {
|
||||||
ban, banned := rq.h.ledger.Check(rq.netblock(), now)
|
ban, banned := rq.h.ledger.Check(rq.client, now)
|
||||||
if banned {
|
if banned {
|
||||||
rq.line.BanExpires = banExpires(ban)
|
rq.line.BanExpires = banExpires(ban)
|
||||||
}
|
}
|
||||||
@@ -36,7 +36,8 @@ func (rq *request) limitBroken(now time.Time) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
ban := rq.h.ledger.BanForLimit(rq.netblock(), now, bans.Notes{
|
netblock := rq.netblock()
|
||||||
|
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
|
||||||
Country: rq.line.Country,
|
Country: rq.line.Country,
|
||||||
Limit: hit.Limit,
|
Limit: hit.Limit,
|
||||||
Window: hit.Window,
|
Window: hit.Window,
|
||||||
@@ -49,6 +50,8 @@ func (rq *request) limitBroken(now time.Time) bool {
|
|||||||
Status: rq.h.config.BanResponse,
|
Status: rq.h.config.BanResponse,
|
||||||
UserAgent: rq.in.UserAgent(),
|
UserAgent: rq.in.UserAgent(),
|
||||||
},
|
},
|
||||||
|
// The histories count this request only once it has ended.
|
||||||
|
Requests: rq.h.limiter.Requests(netblock) + 1,
|
||||||
})
|
})
|
||||||
rq.h.limiter.Reset(group)
|
rq.h.limiter.Reset(group)
|
||||||
|
|
||||||
|
|||||||
@@ -133,7 +133,7 @@ func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) {
|
|||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
|
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
banned := proxy.LedgerOf(server).Bans(netip.MustParsePrefix(client + "/32"))
|
banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
|
||||||
if len(banned) != 2 || banned[0].Notes.Refused != 3 {
|
if len(banned) != 2 || banned[0].Notes.Refused != 3 {
|
||||||
t.Errorf("bans %+v, want two, the first with 3 requests refused", banned)
|
t.Errorf("bans %+v, want two, the first with 3 requests refused", banned)
|
||||||
}
|
}
|
||||||
@@ -291,12 +291,15 @@ func TestBanNotes(t *testing.T) {
|
|||||||
Status: http.StatusForbidden,
|
Status: http.StatusForbidden,
|
||||||
UserAgent: userAgent,
|
UserAgent: userAgent,
|
||||||
},
|
},
|
||||||
|
// The one let through, the one that broke the limit and the two
|
||||||
|
// refused under the ban.
|
||||||
|
Requests: 4,
|
||||||
Refused: 2,
|
Refused: 2,
|
||||||
EarlierBans: 0,
|
EarlierBans: 0,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
ledger := proxy.LedgerOf(server)
|
ledger := server.Ledger
|
||||||
|
|
||||||
got := ledger.Bans(netblock)
|
got := ledger.Bans(netblock)
|
||||||
if len(got) != 1 || got[0] != want {
|
if len(got) != 1 || got[0] != want {
|
||||||
@@ -361,7 +364,7 @@ func (c *clock) advance(d time.Duration) {
|
|||||||
// set to midnight, the start of a bucket in every window.
|
// set to midnight, the start of a bucket in every window.
|
||||||
func startWithClock(
|
func startWithClock(
|
||||||
t *testing.T, geojsURL string, env map[string]string,
|
t *testing.T, geojsURL string, env map[string]string,
|
||||||
) (*sender, *clock, *http.Server) {
|
) (*sender, *clock, *proxy.Server) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
|||||||
@@ -1,15 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/http"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
)
|
|
||||||
|
|
||||||
// LedgerOf returns the ban ledger of a server New returned, so that the
|
|
||||||
// tests can read the bans' notes.
|
|
||||||
func LedgerOf(server *http.Server) *bans.Ledger {
|
|
||||||
h, _ := server.Handler.(*handler)
|
|
||||||
|
|
||||||
return h.ledger
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojsURL, _ := startGeoJS(t)
|
||||||
|
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
||||||
|
rateLimitPerMinute: "2",
|
||||||
|
deniedCountries: "kp",
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
// Two let through, one over the limit, which bans the client, and one
|
||||||
|
// refused under that ban, for which the country is not looked up.
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
clk.advance(time.Second)
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
clk.advance(time.Second)
|
||||||
|
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
want := ratelimit.History{
|
||||||
|
FirstSeen: start,
|
||||||
|
LastSeen: start.Add(2 * time.Second),
|
||||||
|
Country: "DE",
|
||||||
|
LookedUp: start.Add(time.Second),
|
||||||
|
Requests: 4,
|
||||||
|
Forwarded: 2,
|
||||||
|
Refused: 2,
|
||||||
|
// The app answers with no body, smallwebwaf with its status text.
|
||||||
|
ResponseBytes: 2 * int64(len("Forbidden\n")),
|
||||||
|
Responses: ratelimit.Responses{Status2xx: 2, Status4xx: 2},
|
||||||
|
Offences: ratelimit.Offences{Limit: 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := historyOf(t, server, fromDE)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("history\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
_, _ = io.WriteString(w, "hello")
|
||||||
|
})
|
||||||
|
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
|
||||||
|
|
||||||
|
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
|
||||||
|
wantStatus(t, got, http.StatusOK)
|
||||||
|
out.requestLine(t)
|
||||||
|
|
||||||
|
history := historyOf(t, server, localhost)
|
||||||
|
if history.RequestBytes != 3 || history.ResponseBytes != 5 {
|
||||||
|
t.Errorf("history counts %d bytes in and %d out, want 3 and 5",
|
||||||
|
history.RequestBytes, history.ResponseBytes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthEndpointIsNotInTheHistory(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
|
||||||
|
|
||||||
|
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
|
||||||
|
out.requestLine(t)
|
||||||
|
|
||||||
|
if clients := server.Limiter.Snapshot(); len(clients) != 0 {
|
||||||
|
t.Errorf("the table holds %+v, want no client", clients)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// historyOf returns the history of the client at addr.
|
||||||
|
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
client := netip.MustParsePrefix(addr + "/32")
|
||||||
|
for _, c := range server.Limiter.Snapshot() {
|
||||||
|
if c.Client == client {
|
||||||
|
return c.History
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Fatalf("%s is not in the table", client)
|
||||||
|
|
||||||
|
return ratelimit.History{}
|
||||||
|
}
|
||||||
+55
-35
@@ -38,53 +38,70 @@ type Params struct {
|
|||||||
// lookup.URL. GeoJS is asked only while a country list is set.
|
// lookup.URL. GeoJS is asked only while a country list is set.
|
||||||
GeoJSURL string
|
GeoJSURL string
|
||||||
// Now tells the time by which requests are counted for the rate
|
// Now tells the time by which requests are counted for the rate
|
||||||
// limits and bans are made and run out, normally time.Now.
|
// limits, bans are made and run out, and GeoJS's answers are kept,
|
||||||
|
// normally time.Now in UTC, the time the state files give.
|
||||||
Now func() time.Time
|
Now func() time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Server is the server smallwebwaf runs, with the parts of the proxy
|
||||||
|
// whose state the state files keep.
|
||||||
|
type Server struct {
|
||||||
|
*http.Server
|
||||||
|
|
||||||
|
Ledger *bans.Ledger
|
||||||
|
Limiter *ratelimit.Limiter
|
||||||
|
GeoJS *lookup.GeoJS
|
||||||
|
}
|
||||||
|
|
||||||
// New returns the server smallwebwaf runs: each request it reads passes
|
// New returns the server smallwebwaf runs: each request it reads passes
|
||||||
// through the proxy. Go's server itself refuses a request line and
|
// through the proxy. Go's server itself refuses a request line and
|
||||||
// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
|
// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
|
||||||
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
|
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
|
||||||
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
|
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
|
||||||
// applies the timeouts and size limits from then on.
|
// applies the timeouts and size limits from then on.
|
||||||
func New(params Params) *http.Server {
|
func New(params Params) *Server {
|
||||||
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
||||||
|
h := &handler{
|
||||||
|
config: params.Config,
|
||||||
|
requestLog: params.RequestLog,
|
||||||
|
processLog: params.ProcessLog,
|
||||||
|
errorLog: errorLog,
|
||||||
|
transport: newTransport(),
|
||||||
|
now: params.Now,
|
||||||
|
limiter: ratelimit.New(ratelimit.Limits{
|
||||||
|
PerMinute: params.Config.RateLimitPerMinute,
|
||||||
|
PerHour: params.Config.RateLimitPerHour,
|
||||||
|
PerDay: params.Config.RateLimitPerDay,
|
||||||
|
}),
|
||||||
|
ledger: bans.New(bans.Rules{
|
||||||
|
LimitBanDuration: params.Config.LimitBanDuration,
|
||||||
|
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
|
||||||
|
MaxBanDuration: params.Config.MaxBanDuration,
|
||||||
|
MaxBans: params.Config.MaxBans,
|
||||||
|
}),
|
||||||
|
geojs: lookup.New(lookup.Params{
|
||||||
|
URL: params.GeoJSURL,
|
||||||
|
Now: params.Now,
|
||||||
|
ProcessLog: params.ProcessLog,
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
|
||||||
return &http.Server{
|
return &Server{
|
||||||
Addr: params.Config.ListenAddr,
|
Server: &http.Server{
|
||||||
Handler: &handler{
|
Addr: params.Config.ListenAddr,
|
||||||
config: params.Config,
|
Handler: h,
|
||||||
requestLog: params.RequestLog,
|
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
||||||
processLog: params.ProcessLog,
|
// Off is an IdleTimeout of 0, which Go's server replaces with
|
||||||
errorLog: errorLog,
|
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
|
||||||
transport: newTransport(),
|
IdleTimeout: params.Config.ClientIdleTimeout,
|
||||||
now: params.Now,
|
// Go's server reads 4 KiB past MaxHeaderBytes before it
|
||||||
limiter: ratelimit.New(ratelimit.Limits{
|
// refuses, so the limit a client meets is the setting.
|
||||||
PerMinute: params.Config.RateLimitPerMinute,
|
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
||||||
PerHour: params.Config.RateLimitPerHour,
|
ErrorLog: errorLog,
|
||||||
PerDay: params.Config.RateLimitPerDay,
|
|
||||||
}),
|
|
||||||
ledger: bans.New(bans.Rules{
|
|
||||||
LimitBanDuration: params.Config.LimitBanDuration,
|
|
||||||
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
|
|
||||||
MaxBanDuration: params.Config.MaxBanDuration,
|
|
||||||
MaxBans: params.Config.MaxBans,
|
|
||||||
}),
|
|
||||||
geojs: lookup.New(lookup.Params{
|
|
||||||
URL: params.GeoJSURL,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: params.ProcessLog,
|
|
||||||
}),
|
|
||||||
},
|
},
|
||||||
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
Ledger: h.ledger,
|
||||||
// Off is an IdleTimeout of 0, which Go's server replaces with
|
Limiter: h.limiter,
|
||||||
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
|
GeoJS: h.geojs,
|
||||||
IdleTimeout: params.Config.ClientIdleTimeout,
|
|
||||||
// Go's server reads 4 KiB past MaxHeaderBytes before it refuses,
|
|
||||||
// so the limit a client meets is the setting.
|
|
||||||
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
|
||||||
ErrorLog: errorLog,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -130,6 +147,9 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Once the request has ended, before its log line is written.
|
||||||
|
defer rq.addToHistory()
|
||||||
|
|
||||||
refused := rq.check(r.Context())
|
refused := rq.check(r.Context())
|
||||||
if refused != nil {
|
if refused != nil {
|
||||||
rq.answer(*refused)
|
rq.answer(*refused)
|
||||||
|
|||||||
@@ -200,7 +200,7 @@ func startProxyWithGeoJS(
|
|||||||
func startProxyWithClock(
|
func startProxyWithClock(
|
||||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||||
env map[string]string,
|
env map[string]string,
|
||||||
) (string, *output, *http.Server) {
|
) (string, *output, *proxy.Server) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
|
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -318,6 +319,24 @@ func (rq *request) finish() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// addToHistory adds the request, which has ended, to its client's
|
||||||
|
// history.
|
||||||
|
func (rq *request) addToHistory() {
|
||||||
|
var requestBytes int64
|
||||||
|
if rq.body != nil {
|
||||||
|
requestBytes = rq.body.bytes.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
|
||||||
|
Country: rq.line.Country,
|
||||||
|
Forwarded: !rq.upstreamStart.IsZero(),
|
||||||
|
Status: rq.out.status,
|
||||||
|
RequestBytes: requestBytes,
|
||||||
|
ResponseBytes: rq.out.bytes,
|
||||||
|
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// clientRequestDeadline is when the client must have sent its whole
|
// clientRequestDeadline is when the client must have sent its whole
|
||||||
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
||||||
func (rq *request) clientRequestDeadline() time.Time {
|
func (rq *request) clientRequestDeadline() time.Time {
|
||||||
|
|||||||
@@ -0,0 +1,116 @@
|
|||||||
|
package ratelimit_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestHistoryKeepsEveryRequest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
for i, r := range []ratelimit.Request{
|
||||||
|
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
|
||||||
|
{Forwarded: true, Status: 101},
|
||||||
|
{Forwarded: true, Status: 304, RequestBytes: 5},
|
||||||
|
{Country: "FR", Status: 403, ResponseBytes: 10, BrokeLimit: true},
|
||||||
|
{Forwarded: true, Status: 502, ResponseBytes: 12},
|
||||||
|
// Closed without an answer: refused, and no response.
|
||||||
|
{Status: 0},
|
||||||
|
} {
|
||||||
|
limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := ratelimit.History{
|
||||||
|
FirstSeen: start,
|
||||||
|
LastSeen: start.Add(5 * time.Minute),
|
||||||
|
Country: "FR",
|
||||||
|
LookedUp: start.Add(3 * time.Minute),
|
||||||
|
Requests: 6,
|
||||||
|
Forwarded: 4,
|
||||||
|
Refused: 2,
|
||||||
|
RequestBytes: 15,
|
||||||
|
ResponseBytes: 122,
|
||||||
|
Responses: ratelimit.Responses{
|
||||||
|
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 1, Status5xx: 1,
|
||||||
|
},
|
||||||
|
Offences: ratelimit.Offences{Limit: 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := historyOf(t, limiter, client)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("history\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResetKeepsTheHistory(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
for range limit {
|
||||||
|
wantCount(t, limiter, client, start, "")
|
||||||
|
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
||||||
|
}
|
||||||
|
|
||||||
|
limiter.Reset(client)
|
||||||
|
|
||||||
|
if got := historyOf(t, limiter, client).Requests; got != limit {
|
||||||
|
t.Errorf("the history counts %d requests, want %d", got, limit)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
|
||||||
|
for client, requests := range map[string]int{
|
||||||
|
"198.51.100.9/32": 2,
|
||||||
|
"198.51.100.10/32": 3,
|
||||||
|
"192.0.2.1/32": 5,
|
||||||
|
"2001:db8:5::/64": 7,
|
||||||
|
} {
|
||||||
|
for range requests {
|
||||||
|
limiter.AddToHistory(netip.MustParsePrefix(client), midnight(),
|
||||||
|
ratelimit.Request{})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for netblock, want := range map[string]int64{
|
||||||
|
"198.51.100.9/32": 2,
|
||||||
|
"198.51.100.0/24": 5,
|
||||||
|
"2001:db8:5::/64": 7,
|
||||||
|
"203.0.113.0/24": 0,
|
||||||
|
} {
|
||||||
|
got := limiter.Requests(netip.MustParsePrefix(netblock))
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("%s has sent %d requests, want %d", netblock, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// historyOf returns client's history.
|
||||||
|
func historyOf(
|
||||||
|
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix,
|
||||||
|
) ratelimit.History {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for _, c := range limiter.Snapshot() {
|
||||||
|
if c.Client == client {
|
||||||
|
return c.History
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Fatalf("%s is not in the table", client)
|
||||||
|
|
||||||
|
return ratelimit.History{}
|
||||||
|
}
|
||||||
+249
-42
@@ -1,11 +1,15 @@
|
|||||||
// Package ratelimit counts each client's requests over a minute, an hour
|
// Package ratelimit keeps the table of clients: each client's requests
|
||||||
// and a day, as the "Counting method" section of SPEC.md describes, and
|
// counted over a minute, an hour and a day, as the "Counting method"
|
||||||
// tells when a request takes a client over a rate limit. The counts are
|
// section of SPEC.md describes, which tell when a request takes the client
|
||||||
// kept in memory only, for at most 20,000 clients.
|
// over a rate limit, and each client's history since it was first seen.
|
||||||
|
// At most 20,000 clients are kept, in memory, and written to clients.json
|
||||||
|
// and read from it by the state package.
|
||||||
package ratelimit
|
package ratelimit
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -13,7 +17,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// maxClients is how many clients are kept. Past it, the least recently
|
// maxClients is how many clients are kept. Past it, the least recently
|
||||||
// seen client is dropped, and starts afresh if it comes back.
|
// seen client is dropped, with its history, and starts afresh if it comes
|
||||||
|
// back.
|
||||||
const maxClients = 20000
|
const maxClients = 20000
|
||||||
|
|
||||||
const day = 24 * time.Hour
|
const day = 24 * time.Hour
|
||||||
@@ -26,20 +31,94 @@ type Limits struct {
|
|||||||
PerDay int64
|
PerDay int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// Limiter counts each client's requests against the limits. It is safe
|
// Limiter counts each client's requests against the limits, and keeps
|
||||||
// for concurrent use.
|
// its history. It is safe for concurrent use.
|
||||||
type Limiter struct {
|
type Limiter struct {
|
||||||
|
// windows are the minute, the hour and the day, in the order of
|
||||||
|
// Client.buckets.
|
||||||
windows [3]window
|
windows [3]window
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
// clients holds each client's buckets, one pair for each of windows,
|
clients *simplelru.LRU[netip.Prefix, *Client]
|
||||||
// in the same order.
|
}
|
||||||
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
|
|
||||||
|
// Client is a client in the table, as clients.json holds it: its buckets
|
||||||
|
// in each window, and its history.
|
||||||
|
type Client struct {
|
||||||
|
Client netip.Prefix `json:"client"`
|
||||||
|
Minute Buckets `json:"minute"`
|
||||||
|
Hour Buckets `json:"hour"`
|
||||||
|
Day Buckets `json:"day"`
|
||||||
|
History History `json:"history"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Buckets are a client's two buckets in one window: the requests in the
|
||||||
|
// bucket under way, which began at Start, and in the bucket before it.
|
||||||
|
type Buckets struct {
|
||||||
|
Start time.Time `json:"start"`
|
||||||
|
Current int64 `json:"current"`
|
||||||
|
Previous int64 `json:"previous"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// History is what is known of a client since it was first seen.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
|
type History struct {
|
||||||
|
FirstSeen time.Time `json:"first_seen"`
|
||||||
|
LastSeen time.Time `json:"last_seen"`
|
||||||
|
// Country is the client's country as it was last looked up, and
|
||||||
|
// LookedUp when that was; both are empty while it never was.
|
||||||
|
Country string `json:"country,omitempty"`
|
||||||
|
LookedUp time.Time `json:"looked_up,omitzero"`
|
||||||
|
// Requests are all the client's requests: Forwarded those passed to
|
||||||
|
// the app, Refused those refused before anything reached it.
|
||||||
|
Requests int64 `json:"requests"`
|
||||||
|
Forwarded int64 `json:"forwarded"`
|
||||||
|
Refused int64 `json:"refused"`
|
||||||
|
// RequestBytes and ResponseBytes are the body bytes of its requests
|
||||||
|
// and of the responses it was sent.
|
||||||
|
RequestBytes int64 `json:"request_bytes"`
|
||||||
|
ResponseBytes int64 `json:"response_bytes"`
|
||||||
|
Responses Responses `json:"responses,omitzero"`
|
||||||
|
Offences Offences `json:"offences,omitzero"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Responses are the responses a client was sent, by status class;
|
||||||
|
// Status5xx counts every status from 500 up.
|
||||||
|
type Responses struct {
|
||||||
|
Status1xx int64 `json:"1xx,omitempty"`
|
||||||
|
Status2xx int64 `json:"2xx,omitempty"`
|
||||||
|
Status3xx int64 `json:"3xx,omitempty"`
|
||||||
|
Status4xx int64 `json:"4xx,omitempty"`
|
||||||
|
Status5xx int64 `json:"5xx,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Offences are a client's offences, by kind.
|
||||||
|
type Offences struct {
|
||||||
|
// Limit is its requests that broke a rate limit.
|
||||||
|
Limit int64 `json:"limit"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Request is what a client's history keeps of one of its requests.
|
||||||
|
type Request struct {
|
||||||
|
// Country is the client's country, when the request looked it up.
|
||||||
|
Country string
|
||||||
|
// Forwarded is true for a request passed to the app, false for one
|
||||||
|
// refused before anything reached it.
|
||||||
|
Forwarded bool
|
||||||
|
// Status is what the client was sent, 0 if nothing was.
|
||||||
|
Status int
|
||||||
|
// RequestBytes and ResponseBytes are the body bytes of the request
|
||||||
|
// and of its response.
|
||||||
|
RequestBytes int64
|
||||||
|
ResponseBytes int64
|
||||||
|
// BrokeLimit is true for a request that broke a rate limit.
|
||||||
|
BrokeLimit bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns a Limiter for limits, with no client counted yet.
|
// New returns a Limiter for limits, with no client counted yet.
|
||||||
func New(limits Limits) *Limiter {
|
func New(limits Limits) *Limiter {
|
||||||
clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](maxClients, nil)
|
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
@@ -73,16 +152,12 @@ func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
|
|||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
counts, seen := l.clients.Get(client)
|
|
||||||
if !seen {
|
|
||||||
counts = &[3]buckets{}
|
|
||||||
l.clients.Add(client, counts)
|
|
||||||
}
|
|
||||||
|
|
||||||
var hit Hit
|
var hit Hit
|
||||||
|
|
||||||
for i, w := range l.windows {
|
for i, b := range l.get(client).buckets() {
|
||||||
requests := counts[i].add(now, w.length)
|
w := l.windows[i]
|
||||||
|
|
||||||
|
requests := b.add(now, w.length)
|
||||||
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
|
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
|
||||||
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
|
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
|
||||||
}
|
}
|
||||||
@@ -91,12 +166,135 @@ func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
|
|||||||
return hit, hit.Window != ""
|
return hit, hit.Window != ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reset sets client's counts in every window back to zero.
|
// Reset sets client's counts in every window back to zero. Its history
|
||||||
|
// keeps its totals.
|
||||||
func (l *Limiter) Reset(client netip.Prefix) {
|
func (l *Limiter) Reset(client netip.Prefix) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
l.clients.Remove(client)
|
c, seen := l.clients.Peek(client)
|
||||||
|
if seen {
|
||||||
|
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddToHistory adds r, a request from client at now, to the client's
|
||||||
|
// history.
|
||||||
|
func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
h := &l.get(client).History
|
||||||
|
if h.FirstSeen.IsZero() {
|
||||||
|
h.FirstSeen = now
|
||||||
|
}
|
||||||
|
|
||||||
|
h.LastSeen = now
|
||||||
|
|
||||||
|
if r.Country != "" {
|
||||||
|
h.Country = r.Country
|
||||||
|
h.LookedUp = now
|
||||||
|
}
|
||||||
|
|
||||||
|
h.Requests++
|
||||||
|
if r.Forwarded {
|
||||||
|
h.Forwarded++
|
||||||
|
} else {
|
||||||
|
h.Refused++
|
||||||
|
}
|
||||||
|
|
||||||
|
h.RequestBytes += r.RequestBytes
|
||||||
|
h.ResponseBytes += r.ResponseBytes
|
||||||
|
h.Responses.add(r.Status)
|
||||||
|
|
||||||
|
if r.BrokeLimit {
|
||||||
|
h.Offences.Limit++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Requests returns how many requests the clients inside netblock have
|
||||||
|
// sent, as their histories count them.
|
||||||
|
func (l *Limiter) Requests(netblock netip.Prefix) int64 {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
// Most often the netblock is one client.
|
||||||
|
c, seen := l.clients.Peek(netblock)
|
||||||
|
if seen {
|
||||||
|
return c.History.Requests
|
||||||
|
}
|
||||||
|
|
||||||
|
var requests int64
|
||||||
|
|
||||||
|
for _, c := range l.clients.Values() {
|
||||||
|
if netblock.Overlaps(c.Client) {
|
||||||
|
requests += c.History.Requests
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return requests
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot returns every client in the table, sorted by address, as
|
||||||
|
// clients.json lists them.
|
||||||
|
func (l *Limiter) Snapshot() []Client {
|
||||||
|
l.mu.Lock()
|
||||||
|
|
||||||
|
clients := make([]Client, 0, l.clients.Len())
|
||||||
|
for _, c := range l.clients.Values() {
|
||||||
|
clients = append(clients, *c)
|
||||||
|
}
|
||||||
|
|
||||||
|
l.mu.Unlock()
|
||||||
|
|
||||||
|
slices.SortFunc(clients, func(a, b Client) int {
|
||||||
|
return a.Client.Compare(b.Client)
|
||||||
|
})
|
||||||
|
|
||||||
|
return clients
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load puts clients read from clients.json into a table that holds none
|
||||||
|
// yet, in the order they were last seen, so that the least recently seen
|
||||||
|
// is dropped first. Buckets whose time has passed at now are emptied.
|
||||||
|
func (l *Limiter) Load(clients []Client, now time.Time) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
clients = slices.Clone(clients)
|
||||||
|
slices.SortStableFunc(clients, func(a, b Client) int {
|
||||||
|
return a.History.LastSeen.Compare(b.History.LastSeen)
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, c := range clients {
|
||||||
|
for i, b := range c.buckets() {
|
||||||
|
// The window that ends at now covers neither bucket once it
|
||||||
|
// begins after the bucket under way has ended.
|
||||||
|
length := l.windows[i].length
|
||||||
|
if !now.Add(-length).Before(b.Start.Add(length)) {
|
||||||
|
*b = Buckets{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
l.clients.Add(c.Client, &c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// get returns client's entry in the table, a new one if it has none, and
|
||||||
|
// makes it the most recently seen.
|
||||||
|
func (l *Limiter) get(client netip.Prefix) *Client {
|
||||||
|
c, seen := l.clients.Get(client)
|
||||||
|
if !seen {
|
||||||
|
c = &Client{Client: client}
|
||||||
|
l.clients.Add(client, c)
|
||||||
|
}
|
||||||
|
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// buckets returns c's buckets in the minute, the hour and the day.
|
||||||
|
func (c *Client) buckets() [3]*Buckets {
|
||||||
|
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
|
||||||
}
|
}
|
||||||
|
|
||||||
// window is a length of time over which requests are counted, and the
|
// window is a length of time over which requests are counted, and the
|
||||||
@@ -107,14 +305,6 @@ type window struct {
|
|||||||
limit int64
|
limit int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// buckets are a client's two buckets in one window: the requests in the
|
|
||||||
// bucket under way, which began at start, and in the bucket before it.
|
|
||||||
type buckets struct {
|
|
||||||
start time.Time
|
|
||||||
current int64
|
|
||||||
previous int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// add counts a request at now in a window of length, and returns the
|
// add counts a request at now in a window of length, and returns the
|
||||||
// client's requests in the window that ends at now: those in the bucket
|
// client's requests in the window that ends at now: those in the bucket
|
||||||
// under way, and those in the bucket before it weighted by how much of
|
// under way, and those in the bucket before it weighted by how much of
|
||||||
@@ -125,27 +315,44 @@ type buckets struct {
|
|||||||
// bucket. A request dated more than a second before it means the clock
|
// bucket. A request dated more than a second before it means the clock
|
||||||
// was set back, and the buckets start afresh: otherwise the bucket before
|
// was set back, and the buckets start afresh: otherwise the bucket before
|
||||||
// would keep its full weight until the clock caught up.
|
// would keep its full weight until the clock caught up.
|
||||||
func (b *buckets) add(now time.Time, length time.Duration) float64 {
|
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
||||||
if now.Before(b.start.Add(-time.Second)) {
|
if now.Before(b.Start.Add(-time.Second)) {
|
||||||
*b = buckets{}
|
*b = Buckets{}
|
||||||
}
|
}
|
||||||
|
|
||||||
start := now.Truncate(length)
|
start := now.Truncate(length)
|
||||||
if start.After(b.start) {
|
if start.After(b.Start) {
|
||||||
if start.Equal(b.start.Add(length)) {
|
if start.Equal(b.Start.Add(length)) {
|
||||||
b.previous = b.current
|
b.Previous = b.Current
|
||||||
} else {
|
} else {
|
||||||
b.previous = 0
|
b.Previous = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.start = start
|
b.Start = start
|
||||||
b.current = 0
|
b.Current = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.current++
|
b.Current++
|
||||||
|
|
||||||
elapsed := max(now.Sub(b.start), 0)
|
elapsed := max(now.Sub(b.Start), 0)
|
||||||
covered := 1 - float64(elapsed)/float64(length)
|
covered := 1 - float64(elapsed)/float64(length)
|
||||||
|
|
||||||
return float64(b.previous)*covered + float64(b.current)
|
return float64(b.Previous)*covered + float64(b.Current)
|
||||||
|
}
|
||||||
|
|
||||||
|
// add counts a response with status in its class. A status of 0, for
|
||||||
|
// nothing sent, is not a response.
|
||||||
|
func (r *Responses) add(status int) {
|
||||||
|
switch {
|
||||||
|
case status >= http.StatusInternalServerError:
|
||||||
|
r.Status5xx++
|
||||||
|
case status >= http.StatusBadRequest:
|
||||||
|
r.Status4xx++
|
||||||
|
case status >= http.StatusMultipleChoices:
|
||||||
|
r.Status3xx++
|
||||||
|
case status >= http.StatusOK:
|
||||||
|
r.Status2xx++
|
||||||
|
case status >= http.StatusContinue:
|
||||||
|
r.Status1xx++
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package ratelimit_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSnapshotListsTheClientsByAddress(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"}
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
for _, i := range []int{2, 3, 0, 1} {
|
||||||
|
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
|
||||||
|
}
|
||||||
|
|
||||||
|
snapshot := limiter.Snapshot()
|
||||||
|
|
||||||
|
got := make([]string, 0, len(snapshot))
|
||||||
|
for _, c := range snapshot {
|
||||||
|
got = append(got, c.Client.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("snapshot %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
counted := ratelimit.Buckets{Start: midnight(), Current: 1}
|
||||||
|
if snapshot[0].Minute != counted || snapshot[0].Day != counted {
|
||||||
|
t.Errorf("buckets %+v and %+v, want %+v", snapshot[0].Minute, snapshot[0].Day,
|
||||||
|
counted)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadedCountsCarryOn(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
before := ratelimit.New(ratelimit.Limits{PerHour: limit})
|
||||||
|
for range limit {
|
||||||
|
wantCount(t, before, client, start, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Loaded into a new limiter, as across a restart, the client has no
|
||||||
|
// fresh allowance.
|
||||||
|
later := start.Add(time.Minute)
|
||||||
|
after := ratelimit.New(ratelimit.Limits{PerHour: limit})
|
||||||
|
after.Load(before.Snapshot(), later)
|
||||||
|
wantCount(t, after, client, later, hour)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
limiter.Count(client, start)
|
||||||
|
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
||||||
|
|
||||||
|
loaded := func(now time.Time) ratelimit.Client {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
after := ratelimit.New(ratelimit.Limits{})
|
||||||
|
after.Load(limiter.Snapshot(), now)
|
||||||
|
|
||||||
|
return after.Snapshot()[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Two minutes on, the window that ends then covers neither of the
|
||||||
|
// minute's buckets, which are emptied; the hour's and the day's stay,
|
||||||
|
// and so does the history.
|
||||||
|
got := loaded(start.Add(2 * time.Minute))
|
||||||
|
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
|
||||||
|
got.Day.Current != 1 || got.History.Requests != 1 {
|
||||||
|
t.Errorf("loaded two minutes on as %+v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A moment before, the window still covers some of the earlier one.
|
||||||
|
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
|
||||||
|
if got.Minute.Current != 1 {
|
||||||
|
t.Errorf("loaded just under two minutes on with minute buckets %+v",
|
||||||
|
got.Minute)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const maxClients = 20000
|
||||||
|
|
||||||
|
// clients.json lists the clients by address. Here each was last seen
|
||||||
|
// a second before the one listed before it, so the last listed is the
|
||||||
|
// one seen longest ago, and the one dropped.
|
||||||
|
clients := make([]ratelimit.Client, maxClients+1)
|
||||||
|
addr := netip.MustParseAddr("10.0.0.0")
|
||||||
|
|
||||||
|
for i := range clients {
|
||||||
|
clients[i].Client = netip.PrefixFrom(addr, addr.BitLen())
|
||||||
|
clients[i].History.LastSeen = midnight().Add(-time.Duration(i) * time.Second)
|
||||||
|
addr = addr.Next()
|
||||||
|
}
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
limiter.Load(clients, midnight())
|
||||||
|
|
||||||
|
got := limiter.Snapshot()
|
||||||
|
if len(got) != maxClients || got[0].Client != clients[0].Client ||
|
||||||
|
got[maxClients-1].Client != clients[maxClients-1].Client {
|
||||||
|
t.Errorf("%d clients kept, from %s to %s; want %d, from %s to %s",
|
||||||
|
len(got), got[0].Client, got[len(got)-1].Client, maxClients,
|
||||||
|
clients[0].Client, clients[maxClients-1].Client)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -24,12 +24,14 @@ func TestHealthCheck(t *testing.T) {
|
|||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
exited := make(chan int, 1)
|
exited := make(chan int, 1)
|
||||||
|
settings := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: app.URL,
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
exited <- run(ctx, map[string]string{
|
exited <- run(ctx, settings, out)
|
||||||
listenAddr: localhost + ":0",
|
|
||||||
upstreamURL: app.URL,
|
|
||||||
}, out)
|
|
||||||
}()
|
}()
|
||||||
|
|
||||||
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
// Package smallwebwaf runs the smallwebwaf process: it reads the settings
|
||||||
// serves requests until it is told to stop, and then stops in an orderly
|
// and the state files, serves requests until it is told to stop, and then
|
||||||
// way.
|
// stops in an orderly way, writing the state files.
|
||||||
package smallwebwaf
|
package smallwebwaf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -19,6 +19,7 @@ import (
|
|||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/state"
|
||||||
)
|
)
|
||||||
|
|
||||||
// shutdownTimeout is how long requests in progress may take to finish
|
// shutdownTimeout is how long requests in progress may take to finish
|
||||||
@@ -55,8 +56,9 @@ func Main(version string) int {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run reads the settings, then serves requests until ctx is done. It
|
// Run reads the settings and the state files, then serves requests until
|
||||||
// returns the process's exit status, 1 when smallwebwaf cannot start.
|
// ctx is done. It returns the process's exit status, 1 when smallwebwaf
|
||||||
|
// cannot start.
|
||||||
func Run(ctx context.Context, params Params) int {
|
func Run(ctx context.Context, params Params) int {
|
||||||
processLog := requestlog.NewProcessLogger(params.Stdout)
|
processLog := requestlog.NewProcessLogger(params.Stdout)
|
||||||
|
|
||||||
@@ -67,6 +69,33 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The state files give times in UTC.
|
||||||
|
now := func() time.Time { return time.Now().UTC() }
|
||||||
|
|
||||||
|
server := proxy.New(proxy.Params{
|
||||||
|
Config: cfg,
|
||||||
|
RequestLog: params.Stdout,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
GeoJSURL: lookup.URL,
|
||||||
|
Now: now,
|
||||||
|
})
|
||||||
|
|
||||||
|
files, err := state.Load(state.Params{
|
||||||
|
Dir: cfg.StateDir,
|
||||||
|
WriteDelay: cfg.StateWriteDelay,
|
||||||
|
CounterInterval: cfg.StateCounterInterval,
|
||||||
|
Ledger: server.Ledger,
|
||||||
|
Limiter: server.Limiter,
|
||||||
|
GeoJS: server.GeoJS,
|
||||||
|
Now: now,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
processLog.Error("cannot use the state files", "error", err.Error())
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
||||||
@@ -75,27 +104,20 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: params.Stdout,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: lookup.URL,
|
|
||||||
Now: time.Now,
|
|
||||||
})
|
|
||||||
|
|
||||||
processLog.Info("starting",
|
processLog.Info("starting",
|
||||||
"version", params.Version,
|
"version", params.Version,
|
||||||
"address", listener.Addr().String(),
|
"address", listener.Addr().String(),
|
||||||
"settings", cfg)
|
"settings", cfg)
|
||||||
|
|
||||||
return serve(ctx, server, listener, processLog)
|
return serve(ctx, server.Server, listener, files, processLog)
|
||||||
}
|
}
|
||||||
|
|
||||||
// serve serves requests on listener until ctx is done, then gives the
|
// serve serves requests on listener, and writes the state files as they
|
||||||
// requests in progress shutdownTimeout to finish.
|
// are due, until ctx is done. Then it 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,
|
||||||
processLog *slog.Logger,
|
files *state.Files, processLog *slog.Logger,
|
||||||
) int {
|
) int {
|
||||||
served := make(chan error, 1)
|
served := make(chan error, 1)
|
||||||
|
|
||||||
@@ -103,6 +125,16 @@ func serve(
|
|||||||
served <- server.Serve(listener)
|
served <- server.Serve(listener)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
writing, stopWriting := context.WithCancel(ctx)
|
||||||
|
defer stopWriting()
|
||||||
|
|
||||||
|
written := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.Run(writing)
|
||||||
|
close(written)
|
||||||
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-served:
|
case err := <-served:
|
||||||
processLog.Error("serving failed", "error", err.Error())
|
processLog.Error("serving failed", "error", err.Error())
|
||||||
@@ -132,6 +164,21 @@ func serve(
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Run's last write has ended, so nothing else writes the files. Every
|
||||||
|
// request has ended too, but for two kinds that Go's server does not
|
||||||
|
// wait for: one cut off because Shutdown timed out, and one whose
|
||||||
|
// connection switched protocols, such as a WebSocket. Such a request
|
||||||
|
// adds to its client's history only as it ends, which can be after
|
||||||
|
// this write, and then that request is missing from clients.json.
|
||||||
|
<-written
|
||||||
|
|
||||||
|
err = files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
processLog.Error("writing the state files failed", "error", err.Error())
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
processLog.Info("stopped")
|
processLog.Info("stopped")
|
||||||
|
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -24,9 +26,13 @@ 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"
|
||||||
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
|
// greeting is what the tests' app answers.
|
||||||
|
greeting = "hello from the app"
|
||||||
)
|
)
|
||||||
|
|
||||||
// output collects what smallwebwaf writes on stdout.
|
// output collects what smallwebwaf writes on stdout.
|
||||||
@@ -69,11 +75,19 @@ func (o *output) line(t *testing.T, key, value string) map[string]any {
|
|||||||
time.Sleep(pollInterval)
|
time.Sleep(pollInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.buf.String())
|
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.text())
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// text returns everything written so far.
|
||||||
|
func (o *output) text() string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
return o.buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
// run runs smallwebwaf with the settings in env until ctx is done, and
|
// run runs smallwebwaf with the settings in env until ctx is done, and
|
||||||
// returns its exit status.
|
// returns its exit status.
|
||||||
func run(ctx context.Context, env map[string]string, out *output) int {
|
func run(ctx context.Context, env map[string]string, out *output) int {
|
||||||
@@ -121,7 +135,10 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
|
|||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
|
|
||||||
status := run(t.Context(), map[string]string{listenAddr: taken.Addr().String()}, out)
|
status := run(t.Context(), map[string]string{
|
||||||
|
listenAddr: taken.Addr().String(),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
}, out)
|
||||||
if status != 1 {
|
if status != 1 {
|
||||||
t.Errorf("exit status %d, want 1", status)
|
t.Errorf("exit status %d, want 1", status)
|
||||||
}
|
}
|
||||||
@@ -132,11 +149,8 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
|
|||||||
func TestServesUntilToldToStop(t *testing.T) {
|
func TestServesUntilToldToStop(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
app := httptest.NewServer(http.HandlerFunc(
|
appURL := startApp(t)
|
||||||
func(w http.ResponseWriter, _ *http.Request) {
|
dir := t.TempDir()
|
||||||
_, _ = io.WriteString(w, "hello from the app")
|
|
||||||
}))
|
|
||||||
defer app.Close()
|
|
||||||
|
|
||||||
ctx, stop := context.WithCancel(t.Context())
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
out := &output{}
|
out := &output{}
|
||||||
@@ -145,12 +159,13 @@ func TestServesUntilToldToStop(t *testing.T) {
|
|||||||
go func() {
|
go func() {
|
||||||
exited <- run(ctx, map[string]string{
|
exited <- run(ctx, map[string]string{
|
||||||
listenAddr: localhost + ":0",
|
listenAddr: localhost + ":0",
|
||||||
upstreamURL: app.URL,
|
upstreamURL: appURL,
|
||||||
|
stateDir: dir,
|
||||||
}, out)
|
}, out)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
starting := out.line(t, "msg", "starting")
|
starting := out.line(t, "msg", "starting")
|
||||||
wantStartingLine(t, starting, app.URL)
|
wantStartingLine(t, starting, appURL, dir)
|
||||||
|
|
||||||
addr, _ := starting["address"].(string)
|
addr, _ := starting["address"].(string)
|
||||||
wantGreeting(t, "http://"+addr+"/")
|
wantGreeting(t, "http://"+addr+"/")
|
||||||
@@ -170,15 +185,184 @@ func TestServesUntilToldToStop(t *testing.T) {
|
|||||||
out.line(t, "msg", "stopped")
|
out.line(t, "msg", "stopped")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStateKeptAcrossRestarts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rateLimitPerDay: "2",
|
||||||
|
// Neither comes due in the test: the files are written as
|
||||||
|
// smallwebwaf stops.
|
||||||
|
"SWWAF_STATE_WRITE_DELAY": "1h",
|
||||||
|
"SWWAF_STATE_COUNTER_INTERVAL": "1h",
|
||||||
|
}
|
||||||
|
|
||||||
|
// The two requests a day allows, and a stop.
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
wantGreeting(t, url)
|
||||||
|
wantGreeting(t, url)
|
||||||
|
})
|
||||||
|
|
||||||
|
// After a restart the client has no fresh allowance: its third
|
||||||
|
// request breaks the day limit, and bans it.
|
||||||
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
|
wantRefused(t, url)
|
||||||
|
})
|
||||||
|
out.line(t, "action", "rate_limited")
|
||||||
|
|
||||||
|
// After another, the ban still refuses it.
|
||||||
|
out = runUntilStopped(t, env, func(url string) {
|
||||||
|
wantRefused(t, url)
|
||||||
|
})
|
||||||
|
out.line(t, "action", "banned")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const scope = "SWWAF_BAN_SCOPE_V4_PREFIX"
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
"SWWAF_TRUSTED_PROXIES": localhost + "/32",
|
||||||
|
rateLimitPerDay: "1",
|
||||||
|
scope: "24",
|
||||||
|
}
|
||||||
|
|
||||||
|
// 203.0.113.9's second request breaks the day limit, and bans
|
||||||
|
// 203.0.113.0/24.
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
wantStatus(t, url, "203.0.113.9", http.StatusOK)
|
||||||
|
wantStatus(t, url, "203.0.113.9", http.StatusForbidden)
|
||||||
|
})
|
||||||
|
|
||||||
|
// With each address a netblock of its own after a restart, that ban
|
||||||
|
// still refuses all of 203.0.113.0/24. 198.51.100.7 is banned alone.
|
||||||
|
env[scope] = "32"
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
wantStatus(t, url, "203.0.113.200", http.StatusForbidden)
|
||||||
|
wantStatus(t, url, "203.0.114.1", http.StatusOK)
|
||||||
|
wantStatus(t, url, "198.51.100.7", http.StatusOK)
|
||||||
|
wantStatus(t, url, "198.51.100.7", http.StatusForbidden)
|
||||||
|
})
|
||||||
|
|
||||||
|
// With /24 netblocks again, that ban still refuses 198.51.100.7, and
|
||||||
|
// no other address.
|
||||||
|
env[scope] = "24"
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
wantStatus(t, url, "198.51.100.7", http.StatusForbidden)
|
||||||
|
wantStatus(t, url, "198.51.100.8", http.StatusOK)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
err := os.WriteFile(filepath.Join(dir, "bans.json"), []byte("{\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write bans.json: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The file ends at the newline that is the second byte of its first
|
||||||
|
// line.
|
||||||
|
wantStartRefused(t, dir, filepath.Join(dir, "bans.json")+", line 1, column 2: ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnwritableStateDirStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantStartRefused(t, filepath.Join(t.TempDir(), "missing"),
|
||||||
|
"SWWAF_STATE_DIR cannot be written: ")
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantStartRefused runs smallwebwaf with its state files in dir, and
|
||||||
|
// checks that it stops at start, with an error that starts with want. If
|
||||||
|
// it starts instead, it is stopped after waitLimit.
|
||||||
|
func wantStartRefused(t *testing.T, dir, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
out := &output{}
|
||||||
|
|
||||||
|
status := run(ctx, map[string]string{listenAddr: localhost + ":0", stateDir: dir}, out)
|
||||||
|
if status != 1 {
|
||||||
|
t.Fatalf("exit status %d, want 1", status)
|
||||||
|
}
|
||||||
|
|
||||||
|
line := out.line(t, "msg", "cannot use the state files")
|
||||||
|
message, _ := line["error"].(string)
|
||||||
|
|
||||||
|
if !strings.HasPrefix(message, want) {
|
||||||
|
t.Errorf("start refused with %q, want an error starting %q", message, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// startApp starts an app that answers every request with greeting, and
|
||||||
|
// returns its URL.
|
||||||
|
func startApp(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
app := httptest.NewServer(http.HandlerFunc(
|
||||||
|
func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = io.WriteString(w, greeting)
|
||||||
|
}))
|
||||||
|
t.Cleanup(app.Close)
|
||||||
|
|
||||||
|
return app.URL
|
||||||
|
}
|
||||||
|
|
||||||
|
// runUntilStopped runs smallwebwaf with the settings in env, has use send
|
||||||
|
// it requests at url, then stops it as SIGTERM does, checks that it
|
||||||
|
// stopped in order, and returns its output.
|
||||||
|
func runUntilStopped(
|
||||||
|
t *testing.T, env map[string]string, use func(url string),
|
||||||
|
) *output {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
out := &output{}
|
||||||
|
exited := make(chan int, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
exited <- run(ctx, env, out)
|
||||||
|
}()
|
||||||
|
|
||||||
|
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
||||||
|
use("http://" + addr + "/")
|
||||||
|
stop()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case status := <-exited:
|
||||||
|
if status != 0 {
|
||||||
|
t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text())
|
||||||
|
}
|
||||||
|
case <-time.After(waitLimit):
|
||||||
|
t.Fatal("still running after being told to stop")
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
// wantStartingLine checks that the line at start gives the version and
|
// wantStartingLine checks that the line at start gives the version and
|
||||||
// every setting's value.
|
// every setting's value.
|
||||||
func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
|
func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
settings, _ := line["settings"].(map[string]any)
|
settings, _ := line["settings"].(map[string]any)
|
||||||
want := map[string]any{
|
want := map[string]any{
|
||||||
listenAddr: localhost + ":0",
|
listenAddr: localhost + ":0",
|
||||||
upstreamURL: appURL,
|
upstreamURL: appURL,
|
||||||
|
stateDir: dir,
|
||||||
|
"SWWAF_STATE_WRITE_DELAY": "10s",
|
||||||
|
"SWWAF_STATE_COUNTER_INTERVAL": "15m",
|
||||||
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
"SWWAF_TRUSTED_PROXIES": "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",
|
||||||
@@ -193,7 +377,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
|
|||||||
"SWWAF_DENY_NETS": "",
|
"SWWAF_DENY_NETS": "",
|
||||||
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
|
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
|
||||||
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
|
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
|
||||||
"SWWAF_RATE_LIMIT_PER_DAY": "50000",
|
rateLimitPerDay: "50000",
|
||||||
"SWWAF_DENIED_COUNTRIES": "",
|
"SWWAF_DENIED_COUNTRIES": "",
|
||||||
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
|
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
|
||||||
"SWWAF_BAN_RESPONSE": "403",
|
"SWWAF_BAN_RESPONSE": "403",
|
||||||
@@ -236,7 +420,61 @@ func wantGreeting(t *testing.T, url string) {
|
|||||||
body, err := io.ReadAll(res.Body)
|
body, err := io.ReadAll(res.Body)
|
||||||
_ = res.Body.Close()
|
_ = res.Body.Close()
|
||||||
|
|
||||||
if err != nil || string(body) != "hello from the app" {
|
if err != nil || string(body) != greeting {
|
||||||
t.Errorf("got %q (%v), want the app's answer", body, err)
|
t.Errorf("got %q (%v), want the app's answer", body, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// wantRefused checks that a request to url is refused with 403, the
|
||||||
|
// default SWWAF_BAN_RESPONSE.
|
||||||
|
func wantRefused(t *testing.T, url string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
|
||||||
|
http.NoBody)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("new request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
transport := &http.Transport{}
|
||||||
|
defer transport.CloseIdleConnections()
|
||||||
|
|
||||||
|
res, err := (&http.Client{Transport: transport}).Do(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = res.Body.Close()
|
||||||
|
|
||||||
|
if res.StatusCode != http.StatusForbidden {
|
||||||
|
t.Errorf("status %d, want %d", res.StatusCode, http.StatusForbidden)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantStatus checks that a request to url from the client at from, as
|
||||||
|
// X-Forwarded-For names it, is answered with status.
|
||||||
|
func wantStatus(t *testing.T, url, from string, status int) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
|
||||||
|
http.NoBody)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("new request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("X-Forwarded-For", from)
|
||||||
|
|
||||||
|
transport := &http.Transport{}
|
||||||
|
defer transport.CloseIdleConnections()
|
||||||
|
|
||||||
|
res, err := (&http.Client{Transport: transport}).Do(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = res.Body.Close()
|
||||||
|
|
||||||
|
if res.StatusCode != status {
|
||||||
|
t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,492 @@
|
|||||||
|
// Package state keeps smallwebwaf's state in JSON files in
|
||||||
|
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
|
||||||
|
// bans.json holds the bans, clients.json each client's counters and
|
||||||
|
// history, and lookups.json GeoJS's answers. Load reads them at start, and
|
||||||
|
// Run and WriteAll write them, each from a snapshot its part takes under
|
||||||
|
// its own lock, so that no request waits on the disk.
|
||||||
|
package state
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
)
|
||||||
|
|
||||||
|
// version is the version of the files' format, the only one read.
|
||||||
|
const version = 1
|
||||||
|
|
||||||
|
// fileMode lets the smallwebwaf user alone read and write the files, which
|
||||||
|
// hold visitors' addresses.
|
||||||
|
const fileMode = 0o600
|
||||||
|
|
||||||
|
// The state files' names.
|
||||||
|
const (
|
||||||
|
bansJSON = "bans.json"
|
||||||
|
clientsJSON = "clients.json"
|
||||||
|
lookupsJSON = "lookups.json"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errVersion = errors.New("unknown version")
|
||||||
|
// errMissing is for an entry without a field it needs.
|
||||||
|
errMissing = errors.New("has no")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Params are what Load needs.
|
||||||
|
type Params struct {
|
||||||
|
// Dir is the directory of the state files (SWWAF_STATE_DIR).
|
||||||
|
Dir string
|
||||||
|
// WriteDelay is how long after a ban is made bans.json is written
|
||||||
|
// (SWWAF_STATE_WRITE_DELAY), and CounterInterval how often every file
|
||||||
|
// is (SWWAF_STATE_COUNTER_INTERVAL).
|
||||||
|
WriteDelay time.Duration
|
||||||
|
CounterInterval time.Duration
|
||||||
|
// Ledger, Limiter and GeoJS hold the state.
|
||||||
|
Ledger *bans.Ledger
|
||||||
|
Limiter *ratelimit.Limiter
|
||||||
|
GeoJS *lookup.GeoJS
|
||||||
|
// Now tells the time by which the counters' buckets run out, normally
|
||||||
|
// time.Now in UTC.
|
||||||
|
Now func() time.Time
|
||||||
|
// ProcessLog receives what was read, and the writes that fail.
|
||||||
|
ProcessLog *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// Files are the state files of a running smallwebwaf.
|
||||||
|
type Files struct {
|
||||||
|
params Params
|
||||||
|
}
|
||||||
|
|
||||||
|
// bansFile is bans.json, indented for an admin to read and edit.
|
||||||
|
type bansFile struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Bans []banEntry `json:"bans"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// banEntry is a ban as bans.json holds it: a permanent ban's expires is
|
||||||
|
// null.
|
||||||
|
type banEntry struct {
|
||||||
|
Netblock netip.Prefix `json:"netblock"`
|
||||||
|
Start time.Time `json:"start"`
|
||||||
|
Expires *time.Time `json:"expires"`
|
||||||
|
Notes bans.Notes `json:"notes"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// clientsFile is clients.json, with each client on a line of its own.
|
||||||
|
type clientsFile struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Clients []ratelimit.Client `json:"clients"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// lookupsFile is lookups.json, with each answer on a line of its own.
|
||||||
|
type lookupsFile struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
Lookups []lookup.Answer `json:"lookups"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// stateFile is the struct of a state file. Once the file is decoded, its
|
||||||
|
// check refuses the first entry without a field it needs, which would
|
||||||
|
// otherwise be read as something the entry does not say. data is the
|
||||||
|
// file, for a field that may be null or "" but not left out, which the
|
||||||
|
// struct cannot tell apart.
|
||||||
|
type stateFile interface {
|
||||||
|
check(data []byte) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load checks that files can be written in Dir, and reads the state files
|
||||||
|
// in it into the ledger, the limiter and GeoJS. A missing file is empty
|
||||||
|
// state, as on a first start. A file that does not parse, has an unknown
|
||||||
|
// version, or has an entry without a field it needs, is an error that
|
||||||
|
// names the file and, where the JSON decoder tells it, the line and
|
||||||
|
// column, or else the entry.
|
||||||
|
func Load(params Params) (*Files, error) {
|
||||||
|
err := checkWritable(params.Dir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
bansIn bansFile
|
||||||
|
clientsIn clientsFile
|
||||||
|
lookupsIn lookupsFile
|
||||||
|
)
|
||||||
|
|
||||||
|
err = errors.Join(
|
||||||
|
read(params.Dir, bansJSON, &bansIn),
|
||||||
|
read(params.Dir, clientsJSON, &clientsIn),
|
||||||
|
read(params.Dir, lookupsJSON, &lookupsIn),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
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,
|
||||||
|
"bans", len(bansIn.Bans), "clients", len(clientsIn.Clients),
|
||||||
|
"lookups", len(lookupsIn.Lookups))
|
||||||
|
|
||||||
|
return &Files{params: params}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run writes bans.json WriteDelay after a ban is made, with every ban
|
||||||
|
// made in between, and every file every CounterInterval, until ctx is
|
||||||
|
// done. A write that fails is logged, and the file is written again at
|
||||||
|
// its next write.
|
||||||
|
func (f *Files) Run(ctx context.Context) {
|
||||||
|
interval := time.NewTicker(f.params.CounterInterval)
|
||||||
|
defer interval.Stop()
|
||||||
|
|
||||||
|
var bansDue <-chan time.Time // nil while no ban waits to be written
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-f.params.Ledger.Changed():
|
||||||
|
if bansDue == nil {
|
||||||
|
bansDue = time.After(f.params.WriteDelay)
|
||||||
|
}
|
||||||
|
case <-bansDue:
|
||||||
|
bansDue = nil
|
||||||
|
|
||||||
|
f.logFailure(f.writeBans())
|
||||||
|
case <-interval.C:
|
||||||
|
f.logFailure(f.WriteAll())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteAll writes every state file, as smallwebwaf stops. A file that
|
||||||
|
// fails does not keep the others from being written.
|
||||||
|
func (f *Files) WriteAll() error {
|
||||||
|
return errors.Join(f.writeBans(), f.writeClients(), f.writeLookups())
|
||||||
|
}
|
||||||
|
|
||||||
|
// logFailure logs a write that failed.
|
||||||
|
func (f *Files) logFailure(err error) {
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error("writing the state files failed",
|
||||||
|
"error", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeBans writes bans.json.
|
||||||
|
func (f *Files) writeBans() error {
|
||||||
|
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 fmt.Errorf("encode %s: %w", bansJSON, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return write(f.params.Dir, bansJSON, append(data, '\n'))
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeClients writes clients.json.
|
||||||
|
func (f *Files) writeClients() error {
|
||||||
|
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("encode %s: %w", clientsJSON, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return write(f.params.Dir, clientsJSON, data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeLookups writes lookups.json.
|
||||||
|
func (f *Files) writeLookups() error {
|
||||||
|
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return write(f.params.Dir, lookupsJSON, data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// newBanEntry returns ban as bans.json holds it.
|
||||||
|
func newBanEntry(ban bans.Ban) banEntry {
|
||||||
|
entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes}
|
||||||
|
if !ban.Permanent() {
|
||||||
|
entry.Expires = &ban.Expires
|
||||||
|
}
|
||||||
|
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
// ban returns the ban an entry of bans.json holds.
|
||||||
|
func (e banEntry) ban() bans.Ban {
|
||||||
|
ban := bans.Ban{Netblock: e.Netblock, Start: e.Start, Notes: e.Notes}
|
||||||
|
if e.Expires != nil {
|
||||||
|
ban.Expires = *e.Expires
|
||||||
|
}
|
||||||
|
|
||||||
|
return ban
|
||||||
|
}
|
||||||
|
|
||||||
|
// check refuses a ban without a netblock, which would refuse every IPv6
|
||||||
|
// client, a start, from which the length of the netblock's next ban is
|
||||||
|
// worked out, or an expires, which would make it permanent. A permanent
|
||||||
|
// ban's expires is null, which Bans cannot tell from a missing one, so
|
||||||
|
// each expires is read again as written.
|
||||||
|
func (f *bansFile) check(data []byte) error {
|
||||||
|
var written struct {
|
||||||
|
Bans []struct {
|
||||||
|
Expires json.RawMessage `json:"expires"`
|
||||||
|
} `json:"bans"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err := json.Unmarshal(data, &written)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, entry := range f.Bans {
|
||||||
|
switch {
|
||||||
|
case !entry.Netblock.IsValid():
|
||||||
|
return missing(i, "netblock")
|
||||||
|
case entry.Start.IsZero():
|
||||||
|
return missing(i, "start")
|
||||||
|
case written.Bans[i].Expires == nil:
|
||||||
|
return missing(i, "expires")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// check refuses a client without its address, which would count nobody's
|
||||||
|
// requests, or with requests in a window but no start, which would drop
|
||||||
|
// them and give the client a fresh allowance.
|
||||||
|
func (f *clientsFile) check([]byte) error {
|
||||||
|
for i, client := range f.Clients {
|
||||||
|
switch {
|
||||||
|
case !client.Client.IsValid():
|
||||||
|
return missing(i, "client")
|
||||||
|
case countsWithoutStart(client.Minute):
|
||||||
|
return missing(i, "minute.start")
|
||||||
|
case countsWithoutStart(client.Hour):
|
||||||
|
return missing(i, "hour.start")
|
||||||
|
case countsWithoutStart(client.Day):
|
||||||
|
return missing(i, "day.start")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// check refuses an answer without a client, which would answer for
|
||||||
|
// nobody, a country, which would place the client nowhere, or the time
|
||||||
|
// GeoJS gave it, which would drop it. "" is the country of a client
|
||||||
|
// GeoJS cannot place, which Lookups cannot tell from a missing one, so
|
||||||
|
// each country is read again as written.
|
||||||
|
func (f *lookupsFile) check(data []byte) error {
|
||||||
|
var written struct {
|
||||||
|
Lookups []struct {
|
||||||
|
Country *string `json:"country"`
|
||||||
|
} `json:"lookups"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err := json.Unmarshal(data, &written)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, answer := range f.Lookups {
|
||||||
|
switch {
|
||||||
|
case !answer.Client.IsValid():
|
||||||
|
return missing(i, "client")
|
||||||
|
case written.Lookups[i].Country == nil:
|
||||||
|
return missing(i, "country")
|
||||||
|
case answer.Answered.IsZero():
|
||||||
|
return missing(i, "answered")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// countsWithoutStart reports whether b holds requests but no start, which
|
||||||
|
// places them in time.
|
||||||
|
func countsWithoutStart(b ratelimit.Buckets) bool {
|
||||||
|
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// missing returns the error for entry i, counted from 0, of a state file,
|
||||||
|
// which has no field.
|
||||||
|
func missing(i int, field string) error {
|
||||||
|
return fmt.Errorf("entry %d %w %q", i+1, errMissing, field)
|
||||||
|
}
|
||||||
|
|
||||||
|
// encodeOnePerLine encodes a state file whose entries, under key, are one
|
||||||
|
// to a line, so that grep shows everything about one client.
|
||||||
|
func encodeOnePerLine[E any](key string, entries []E) ([]byte, error) {
|
||||||
|
var b bytes.Buffer
|
||||||
|
|
||||||
|
fmt.Fprintf(&b, "{\n \"version\": %d,\n %q: [", version, key)
|
||||||
|
|
||||||
|
for i, entry := range entries {
|
||||||
|
line, err := json.Marshal(entry)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if i > 0 {
|
||||||
|
b.WriteString(",")
|
||||||
|
}
|
||||||
|
|
||||||
|
b.WriteString("\n ")
|
||||||
|
b.Write(line)
|
||||||
|
}
|
||||||
|
|
||||||
|
b.WriteString("\n ]\n}\n")
|
||||||
|
|
||||||
|
return b.Bytes(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkWritable makes a file in dir and removes it again.
|
||||||
|
func checkWritable(dir string) error {
|
||||||
|
file, err := os.CreateTemp(dir, "write-check-*")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Join(file.Close(), os.Remove(file.Name()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// read reads the state file name in dir into file, a pointer to that
|
||||||
|
// file's struct, and checks its entries. A missing file leaves file as it
|
||||||
|
// is.
|
||||||
|
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
|
||||||
|
// refused for that, and not for an entry this version cannot read.
|
||||||
|
var header struct {
|
||||||
|
Version int `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err = json.Unmarshal(data, &header)
|
||||||
|
if err == nil && header.Version != version {
|
||||||
|
err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d",
|
||||||
|
errVersion, header.Version, version)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||||
|
// A field this version does not know is most likely misspelt, and
|
||||||
|
// its value would be lost without a word.
|
||||||
|
decoder.DisallowUnknownFields()
|
||||||
|
err = decoder.Decode(file)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
err = file.check(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s%s: %w", path, position(data, err), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// position returns where in data err was found, as ", line L, column C"
|
||||||
|
// of the last byte the JSON decoder read, or "" when err does not tell.
|
||||||
|
func position(data []byte, err error) string {
|
||||||
|
var (
|
||||||
|
syntaxErr *json.SyntaxError
|
||||||
|
typeErr *json.UnmarshalTypeError
|
||||||
|
read int64
|
||||||
|
)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case errors.As(err, &syntaxErr):
|
||||||
|
read = syntaxErr.Offset
|
||||||
|
case errors.As(err, &typeErr):
|
||||||
|
read = typeErr.Offset
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
before := data[:max(min(read, int64(len(data)))-1, 0)]
|
||||||
|
line := bytes.Count(before, []byte("\n")) + 1
|
||||||
|
column := len(before) - bytes.LastIndexByte(before, '\n')
|
||||||
|
|
||||||
|
return fmt.Sprintf(", line %d, column %d", line, column)
|
||||||
|
}
|
||||||
|
|
||||||
|
// write writes data to the file name in dir so that a crash at any
|
||||||
|
// moment leaves either the old file or the new one, whole: data goes to a
|
||||||
|
// temporary file in the same directory, which is synced and renamed over
|
||||||
|
// name, and then the directory is synced, so that the rename lasts.
|
||||||
|
func write(dir, name string, data []byte) error {
|
||||||
|
path := filepath.Join(dir, name)
|
||||||
|
temporary := path + ".tmp"
|
||||||
|
|
||||||
|
err := writeSynced(temporary, data)
|
||||||
|
if err == nil {
|
||||||
|
err = os.Rename(temporary, path)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(temporary)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Join(directory.Sync(), directory.Close())
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeSynced writes data to the file at path, and syncs it to the disk.
|
||||||
|
func writeSynced(path string, data []byte) error {
|
||||||
|
//nolint:gosec // a state file's temporary file, in SWWAF_STATE_DIR
|
||||||
|
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = file.Write(data)
|
||||||
|
if err == nil {
|
||||||
|
err = file.Sync()
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Join(err, file.Close())
|
||||||
|
}
|
||||||
@@ -0,0 +1,635 @@
|
|||||||
|
package state_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/state"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// The state files.
|
||||||
|
bansJSON = "bans.json"
|
||||||
|
clientsJSON = "clients.json"
|
||||||
|
lookupsJSON = "lookups.json"
|
||||||
|
)
|
||||||
|
|
||||||
|
// permanentBansJSON is bans.json holding permanentBan.
|
||||||
|
const permanentBansJSON = `{
|
||||||
|
"version": 1,
|
||||||
|
"bans": [
|
||||||
|
{
|
||||||
|
"netblock": "2001:db8::/64",
|
||||||
|
"start": "2026-10-06T00:00:00Z",
|
||||||
|
"expires": null,
|
||||||
|
"notes": {
|
||||||
|
"country": "DE",
|
||||||
|
"limit": 1000,
|
||||||
|
"window": "minute",
|
||||||
|
"count": 1000.5,
|
||||||
|
"request": {
|
||||||
|
"time": "2026-10-06T00:00:00Z",
|
||||||
|
"method": "GET",
|
||||||
|
"host": "app.example",
|
||||||
|
"path": "/repo?page=2",
|
||||||
|
"status": 403,
|
||||||
|
"user_agent": "scraper/1.0"
|
||||||
|
},
|
||||||
|
"requests": 1500,
|
||||||
|
"refused": 3,
|
||||||
|
"earlier_bans": 5
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
`
|
||||||
|
|
||||||
|
func TestFilesWrittenAndReadBack(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
before := newParams(dir)
|
||||||
|
fill(before)
|
||||||
|
|
||||||
|
files, err := state.Load(before)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read into new parts, as at the next start, the files give back what
|
||||||
|
// was written.
|
||||||
|
after := newParams(dir)
|
||||||
|
load(t, after)
|
||||||
|
|
||||||
|
wantEqual(t, bansJSON, after.Ledger.Snapshot(), before.Ledger.Snapshot())
|
||||||
|
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
|
||||||
|
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
|
||||||
|
|
||||||
|
// Each one-per-line file lists its entries by client, and nothing
|
||||||
|
// but the three files is left in the directory.
|
||||||
|
wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
|
||||||
|
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
|
||||||
|
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
|
||||||
|
"192.0.2.1/32", "203.0.113.9/32")
|
||||||
|
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
params := newParams(dir)
|
||||||
|
params.Ledger.Load([]bans.Ban{permanentBan()})
|
||||||
|
|
||||||
|
files := load(t, params)
|
||||||
|
|
||||||
|
err := files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := readFile(t, filepath.Join(dir, bansJSON))
|
||||||
|
if got != permanentBansJSON {
|
||||||
|
t.Errorf("bans.json\n%s\nwant\n%s", got, permanentBansJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMissingFilesAreEmptyState(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
params := newParams(t.TempDir())
|
||||||
|
load(t, params)
|
||||||
|
|
||||||
|
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
|
||||||
|
len(params.GeoJS.Snapshot()) != 0 {
|
||||||
|
t.Error("state from no files")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, file, content string
|
||||||
|
// want is what the error says after the file's path.
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"a syntax error", bansJSON,
|
||||||
|
"{\n \"version\": 1,\n \"bans\": [\n" +
|
||||||
|
" {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n",
|
||||||
|
", line 4, column 39: invalid character '}'",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a value of the wrong kind", clientsJSON,
|
||||||
|
"{\n \"version\": 1,\n \"clients\": [\n" +
|
||||||
|
" {\"client\":\"203.0.113.9/32\",\"history\":{\"requests\":\"many\"}}\n" +
|
||||||
|
" ]\n}\n",
|
||||||
|
", line 4, column ",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Found at the newline that ends the file.
|
||||||
|
"a cut-off file", lookupsJSON,
|
||||||
|
"{\n \"version\": 1,\n \"lookups\": [\n",
|
||||||
|
", line 3, column 17: unexpected end of JSON input",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an unknown field", lookupsJSON,
|
||||||
|
`{"version": 1, "lookups": [{"client": "203.0.113.9/32", "contry": "DE"}]}`,
|
||||||
|
`: json: unknown field "contry"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a netblock that does not read", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
|
||||||
|
`: netip.ParsePrefix("203.0.113.300/32")`,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantRefused(t, tc.file, tc.content, tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
// The other fields each entry needs.
|
||||||
|
ban = `"start": "2026-10-06T00:00:00Z", "expires": null`
|
||||||
|
answer = `"answered": "2026-10-06T00:00:00Z"`
|
||||||
|
|
||||||
|
noNetblock = `: entry 1 has no "netblock"`
|
||||||
|
)
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, file, content string
|
||||||
|
// want is what the error says after the file's path.
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"a ban without a netblock", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{` + ban + `}]}`,
|
||||||
|
noNetblock,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a ban whose netblock is null", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{"netblock": null, ` + ban + `}]}`,
|
||||||
|
noNetblock,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a ban whose netblock is empty", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{"netblock": "", ` + ban + `}]}`,
|
||||||
|
noNetblock,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a ban without a start", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{"netblock": "203.0.113.9/32", "expires": null}]}`,
|
||||||
|
`: entry 1 has no "start"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// The first ban's expires is null, as a permanent ban's is.
|
||||||
|
"a ban without an expires", bansJSON,
|
||||||
|
`{"version": 1, "bans": [{"netblock": "203.0.113.9/32", ` + ban + `}, ` +
|
||||||
|
`{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z"}]}`,
|
||||||
|
`: entry 2 has no "expires"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a client without its address", clientsJSON,
|
||||||
|
`{"version": 1, "clients": [{"history": {"requests": 3}}]}`,
|
||||||
|
`: entry 1 has no "client"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a client with requests in a window without its start", clientsJSON,
|
||||||
|
`{"version": 1, "clients": [{"client": "203.0.113.9/32", ` +
|
||||||
|
`"hour": {"current": 3}}]}`,
|
||||||
|
`: entry 1 has no "hour.start"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an answer without a client", lookupsJSON,
|
||||||
|
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
|
||||||
|
`: entry 1 has no "client"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// A country of "" is a client GeoJS cannot place.
|
||||||
|
"an answer without a country", lookupsJSON,
|
||||||
|
`{"version": 1, "lookups": [{"client": "192.0.2.1/32", "country": "", ` +
|
||||||
|
answer + `}, {"client": "203.0.113.9/32", ` + answer + `}]}`,
|
||||||
|
`: entry 2 has no "country"`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an answer without the time GeoJS gave it", lookupsJSON,
|
||||||
|
`{"version": 1, "lookups": [{"client": "203.0.113.9/32", "country": "DE"}]}`,
|
||||||
|
`: entry 1 has no "answered"`,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantRefused(t, tc.file, tc.content, tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnknownVersionStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} {
|
||||||
|
for _, content := range []string{`{"version": 2}`, `{}`} {
|
||||||
|
t.Run(file+" "+content, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantRefused(t, file, content, ": unknown version ")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnwritableDirectoryStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
notADirectory := filepath.Join(t.TempDir(), "file")
|
||||||
|
|
||||||
|
err := os.WriteFile(notADirectory, nil, 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, dir := range []string{
|
||||||
|
filepath.Join(t.TempDir(), "missing"),
|
||||||
|
notADirectory,
|
||||||
|
} {
|
||||||
|
const want = "SWWAF_STATE_DIR cannot be written: "
|
||||||
|
|
||||||
|
_, err := state.Load(newParams(dir))
|
||||||
|
if err == nil || !strings.HasPrefix(err.Error(), want) {
|
||||||
|
t.Errorf("state directory %s: error %v, want one starting %s", dir, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The two 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
|
||||||
|
// returns once Run waits for its next write, so that every write due by
|
||||||
|
// then is on disk.
|
||||||
|
|
||||||
|
func TestBansWrittenOnceWriteDelayAfterABan(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))
|
||||||
|
|
||||||
|
// A second ban, made while the first waits to be written, puts the
|
||||||
|
// write off no further, and is written with it.
|
||||||
|
first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
|
||||||
|
midnight(), bans.Notes{})
|
||||||
|
|
||||||
|
time.Sleep(5 * time.Second)
|
||||||
|
|
||||||
|
second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
|
||||||
|
midnight(), bans.Notes{})
|
||||||
|
|
||||||
|
time.Sleep(5*time.Second - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFiles(t, dir)
|
||||||
|
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFiles(t, dir, bansJSON)
|
||||||
|
|
||||||
|
read := newParams(dir)
|
||||||
|
load(t, read)
|
||||||
|
|
||||||
|
want := []bans.Ban{first, second}
|
||||||
|
if got := read.Ledger.Snapshot(); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("bans.json holds %+v, want %+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// That write was the only one: bans.json is not written again for
|
||||||
|
// the second ban. The other files wait for the interval, an hour
|
||||||
|
// away.
|
||||||
|
removeFiles(t, dir, bansJSON)
|
||||||
|
time.Sleep(params.WriteDelay)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFiles(t, dir)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
params := newParams(dir)
|
||||||
|
params.CounterInterval = time.Minute
|
||||||
|
run(t, load(t, params))
|
||||||
|
|
||||||
|
// The files are removed once written, so that each interval shows
|
||||||
|
// them written again.
|
||||||
|
for range 3 {
|
||||||
|
time.Sleep(time.Minute - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFiles(t, dir)
|
||||||
|
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
params := newParams(dir)
|
||||||
|
params.Ledger.Load([]bans.Ban{permanentBan()})
|
||||||
|
files := load(t, params)
|
||||||
|
|
||||||
|
err := files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A directory in the way of bans.json's temporary file fails its
|
||||||
|
// next write, but not the others'.
|
||||||
|
err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("mkdir: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
||||||
|
bans.Notes{})
|
||||||
|
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
|
||||||
|
|
||||||
|
err = files.WriteAll()
|
||||||
|
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
|
||||||
|
t.Errorf("error %v, want one naming bans.json's temporary file", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := readFile(t, filepath.Join(dir, bansJSON))
|
||||||
|
if got != permanentBansJSON {
|
||||||
|
t.Errorf("bans.json is now\n%s\nwant it as it was", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
read := newParams(dir)
|
||||||
|
load(t, read)
|
||||||
|
|
||||||
|
if len(read.Limiter.Snapshot()) != 1 {
|
||||||
|
t.Error("clients.json was not written")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
files := load(t, newParams(dir))
|
||||||
|
|
||||||
|
// 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 = files.WriteAll()
|
||||||
|
if err == nil {
|
||||||
|
t.Error("writing over a directory did not fail")
|
||||||
|
}
|
||||||
|
|
||||||
|
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
||||||
|
}
|
||||||
|
|
||||||
|
// midnight is the time of the tests' clock.
|
||||||
|
func midnight() time.Time {
|
||||||
|
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
|
||||||
|
}
|
||||||
|
|
||||||
|
// newParams returns Params for the state files in dir, with parts that
|
||||||
|
// hold nothing yet. GeoJS is never asked.
|
||||||
|
func newParams(dir string) state.Params {
|
||||||
|
discard := slog.New(slog.DiscardHandler)
|
||||||
|
|
||||||
|
return state.Params{
|
||||||
|
Dir: dir,
|
||||||
|
WriteDelay: time.Hour,
|
||||||
|
CounterInterval: time.Hour,
|
||||||
|
Ledger: bans.New(bans.Rules{
|
||||||
|
LimitBanDuration: time.Hour,
|
||||||
|
LimitBanRepeatWindow: 24 * time.Hour,
|
||||||
|
MaxBanDuration: 7 * 24 * time.Hour,
|
||||||
|
MaxBans: 5000,
|
||||||
|
}),
|
||||||
|
Limiter: ratelimit.New(ratelimit.Limits{}),
|
||||||
|
GeoJS: lookup.New(lookup.Params{Now: midnight, ProcessLog: discard}),
|
||||||
|
Now: midnight,
|
||||||
|
ProcessLog: discard,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// fill puts a ban that ends and one that does not, clients with counts
|
||||||
|
// and histories, and GeoJS answers into the parts of params.
|
||||||
|
func fill(params state.Params) {
|
||||||
|
now := midnight()
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
params.Ledger.Load([]bans.Ban{permanentBan()})
|
||||||
|
params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1})
|
||||||
|
|
||||||
|
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
|
||||||
|
params.Limiter.Count(netip.MustParsePrefix(c), now)
|
||||||
|
}
|
||||||
|
|
||||||
|
params.Limiter.AddToHistory(client, now, ratelimit.Request{
|
||||||
|
Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
|
||||||
|
})
|
||||||
|
|
||||||
|
params.GeoJS.Load([]lookup.Answer{
|
||||||
|
{Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now},
|
||||||
|
{
|
||||||
|
Client: netip.MustParsePrefix("192.0.2.1/32"),
|
||||||
|
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// permanentBan is the ban permanentBansJSON holds.
|
||||||
|
func permanentBan() bans.Ban {
|
||||||
|
return bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("2001:db8::/64"),
|
||||||
|
Start: midnight(),
|
||||||
|
Notes: bans.Notes{
|
||||||
|
Country: "DE",
|
||||||
|
Limit: 1000,
|
||||||
|
Window: "minute",
|
||||||
|
Count: 1000.5,
|
||||||
|
Request: bans.Request{
|
||||||
|
Time: midnight(),
|
||||||
|
Method: "GET",
|
||||||
|
Host: "app.example",
|
||||||
|
Path: "/repo?page=2",
|
||||||
|
Status: 403,
|
||||||
|
UserAgent: "scraper/1.0",
|
||||||
|
},
|
||||||
|
Requests: 1500,
|
||||||
|
Refused: 3,
|
||||||
|
EarlierBans: 5,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// load reads the state files into the parts of params.
|
||||||
|
func load(t *testing.T, params state.Params) *state.Files {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
files, err := state.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return files
|
||||||
|
}
|
||||||
|
|
||||||
|
// run runs files' writes until the test ends.
|
||||||
|
func run(t *testing.T, files *state.Files) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.Run(ctx)
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
stop()
|
||||||
|
<-stopped
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantEqual checks that the entries read back from file are those
|
||||||
|
// written.
|
||||||
|
func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// readFile returns what the file at path holds.
|
||||||
|
func readFile(t *testing.T, path string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path) //nolint:gosec // a file the test wrote
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return string(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRefused writes content to the state file named file in a new
|
||||||
|
// directory, and checks that Load refuses it with an error that is the
|
||||||
|
// file's path and then starts with want.
|
||||||
|
func wantRefused(t *testing.T, file, content, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, file)
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", file, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = state.Load(newParams(dir))
|
||||||
|
if err == nil || !strings.HasPrefix(err.Error(), path+want) {
|
||||||
|
t.Errorf("error %v, want one starting %s%s", err, path, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeFiles removes the named files from dir.
|
||||||
|
func removeFiles(t *testing.T, dir string, names ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
|
err := os.Remove(filepath.Join(dir, name))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("remove: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantFiles checks the names of the files in dir.
|
||||||
|
func wantFiles(t *testing.T, dir string, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
entries, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read %s: %v", dir, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := make([]string, 0, len(entries))
|
||||||
|
for _, entry := range entries {
|
||||||
|
got = append(got, entry.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("%s holds %v, want %v", dir, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantEntries checks that the file at path has its version, then its
|
||||||
|
// entries under key, each on a line of its own, for the clients want
|
||||||
|
// names in that order.
|
||||||
|
func wantEntries(t *testing.T, path, key string, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
data := readFile(t, path)
|
||||||
|
lines := strings.Split(strings.TrimSuffix(data, "\n"), "\n")
|
||||||
|
head := []string{"{", ` "version": 1,`, ` "` + key + `": [`}
|
||||||
|
tail := []string{" ]", "}"}
|
||||||
|
|
||||||
|
if len(lines) != len(head)+len(want)+len(tail) ||
|
||||||
|
!slices.Equal(lines[:len(head)], head) ||
|
||||||
|
!slices.Equal(lines[len(lines)-len(tail):], tail) {
|
||||||
|
t.Fatalf("%s is\n%s", path, data)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, client := range want {
|
||||||
|
line := strings.TrimSuffix(lines[len(head)+i], ",")
|
||||||
|
|
||||||
|
var entry struct {
|
||||||
|
Client string `json:"client"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err := json.Unmarshal([]byte(line), &entry)
|
||||||
|
if err != nil || entry.Client != client {
|
||||||
|
t.Errorf("entry %d of %s is %s (%v), want %s's", i, path, line, err, client)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+44
-11
@@ -1,12 +1,14 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# script/example-app: build the image, and on it the example app in
|
# script/example-app: build the image, and on it the example app in
|
||||||
# deploy/example-app, then run the app's container and check that the
|
# deploy/example-app, then run the app's container with a volume for the
|
||||||
# health check passes, that a request is served through smallwebwaf,
|
# state files and check that the health check passes, that a request is
|
||||||
# that `sv stop` stops smallwebwaf in order, and that `docker stop`
|
# served through smallwebwaf, that a second one in a minute bans the
|
||||||
# stops the container without having to kill it. The container and both
|
# client, that `sv stop` stops smallwebwaf in order, that `docker stop`
|
||||||
# images are removed however the script ends. Building the app needs
|
# stops the container without having to kill it, and that a new
|
||||||
# network access, for nixpkgs' binary cache. script/check does not run
|
# container on the same volume still refuses the banned client. The
|
||||||
# this.
|
# containers, the volume and both images are removed however the script
|
||||||
|
# ends. Building the app needs network access, for nixpkgs' binary cache.
|
||||||
|
# script/check does not run this.
|
||||||
set -eu
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -18,9 +20,11 @@ NAME="$("$SCRIPT_DIR/projectname")-example-$$"
|
|||||||
IMAGE="$NAME-base"
|
IMAGE="$NAME-base"
|
||||||
APP_IMAGE="$NAME-app"
|
APP_IMAGE="$NAME-app"
|
||||||
CONTAINER="$NAME"
|
CONTAINER="$NAME"
|
||||||
|
VOLUME="$NAME-state"
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
docker rm --force "$CONTAINER" >/dev/null 2>&1 || true
|
docker rm --force "$CONTAINER" >/dev/null 2>&1 || true
|
||||||
|
docker volume rm --force "$VOLUME" >/dev/null 2>&1 || true
|
||||||
docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true
|
docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -53,6 +57,26 @@ logged() {
|
|||||||
docker logs "$CONTAINER" 2>&1 | grep -qF "$1"
|
docker logs "$CONTAINER" 2>&1 | grep -qF "$1"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# start_container: run the app's container, with the state files on the
|
||||||
|
# volume and a rate limit of one request a minute, and wait until it is
|
||||||
|
# healthy.
|
||||||
|
start_container() {
|
||||||
|
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
||||||
|
--volume "$VOLUME:/var/lib/smallwebwaf" \
|
||||||
|
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
|
||||||
|
"$APP_IMAGE" >/dev/null
|
||||||
|
wait_for "the health check did not pass" healthy
|
||||||
|
address="$(docker port "$CONTAINER" 8080/tcp)"
|
||||||
|
}
|
||||||
|
|
||||||
|
# refused: a request to the container gets 403, SWWAF_BAN_RESPONSE's
|
||||||
|
# default.
|
||||||
|
refused() {
|
||||||
|
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
|
||||||
|
--max-time 10 "http://$address/")" || true
|
||||||
|
[ "$code" = 403 ]
|
||||||
|
}
|
||||||
|
|
||||||
main() {
|
main() {
|
||||||
cd "$ROOT"
|
cd "$ROOT"
|
||||||
trap cleanup EXIT
|
trap cleanup EXIT
|
||||||
@@ -62,18 +86,20 @@ main() {
|
|||||||
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
|
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
|
||||||
-t "$APP_IMAGE" deploy/example-app
|
-t "$APP_IMAGE" deploy/example-app
|
||||||
|
|
||||||
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
docker volume create "$VOLUME" >/dev/null
|
||||||
"$APP_IMAGE" >/dev/null
|
start_container
|
||||||
wait_for "the health check did not pass" healthy
|
|
||||||
echo "example-app: the health check passes"
|
echo "example-app: the health check passes"
|
||||||
|
|
||||||
address="$(docker port "$CONTAINER" 8080/tcp)"
|
|
||||||
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
|
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
|
||||||
fail "no answer on port 8080"
|
fail "no answer on port 8080"
|
||||||
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
|
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
|
||||||
wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"'
|
wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"'
|
||||||
echo "example-app: smallwebwaf passes a request to the app and its answer back"
|
echo "example-app: smallwebwaf passes a request to the app and its answer back"
|
||||||
|
|
||||||
|
refused || fail "a second request in a minute was not refused"
|
||||||
|
wait_for "smallwebwaf logged no ban" logged '"action":"rate_limited"'
|
||||||
|
echo "example-app: a second request in a minute bans the client"
|
||||||
|
|
||||||
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
|
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
|
||||||
fail "sv stop smallwebwaf failed"
|
fail "sv stop smallwebwaf failed"
|
||||||
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
|
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
|
||||||
@@ -83,6 +109,13 @@ main() {
|
|||||||
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
|
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
|
||||||
[ "$status" = 0 ] || fail "docker stop left exit status $status"
|
[ "$status" = 0 ] || fail "docker stop left exit status $status"
|
||||||
echo "example-app: docker stop stops the container in order"
|
echo "example-app: docker stop stops the container in order"
|
||||||
|
|
||||||
|
docker rm "$CONTAINER" >/dev/null
|
||||||
|
start_container
|
||||||
|
refused || fail "the new container let the banned client through"
|
||||||
|
wait_for "smallwebwaf logged no request refused under the ban" \
|
||||||
|
logged '"action":"banned"'
|
||||||
|
echo "example-app: a new container on the same volume keeps the ban"
|
||||||
}
|
}
|
||||||
|
|
||||||
main "$@"
|
main "$@"
|
||||||
|
|||||||
+7
-1
@@ -1,6 +1,7 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# script/run: build bin/smallwebwaf with script/build and run it, with
|
# script/run: build bin/smallwebwaf with script/build and run it, with
|
||||||
# the settings in the environment.
|
# the settings in the environment. Unless SWWAF_STATE_DIR is set, the
|
||||||
|
# state files go in bin/state, beside the binary.
|
||||||
set -eu
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -8,6 +9,11 @@ ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)"
|
|||||||
|
|
||||||
main() {
|
main() {
|
||||||
"$SCRIPT_DIR/build"
|
"$SCRIPT_DIR/build"
|
||||||
|
if [ -z "${SWWAF_STATE_DIR+set}" ]; then
|
||||||
|
SWWAF_STATE_DIR="$ROOT/bin/state"
|
||||||
|
export SWWAF_STATE_DIR
|
||||||
|
mkdir -p "$SWWAF_STATE_DIR"
|
||||||
|
fi
|
||||||
exec "$ROOT/bin/smallwebwaf"
|
exec "$ROOT/bin/smallwebwaf"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -2,10 +2,14 @@
|
|||||||
set -euo pipefail
|
set -euo pipefail
|
||||||
|
|
||||||
# runit's run script for smallwebwaf, run again whenever smallwebwaf
|
# runit's run script for smallwebwaf, run again whenever smallwebwaf
|
||||||
# exits; the wait spaces out the restarts. exec, so that the signal
|
# exits; the wait spaces out the restarts. The state directory and every
|
||||||
# `sv stop` sends reaches smallwebwaf itself.
|
# file in it are given to the smallwebwaf user, so that a volume mounted
|
||||||
|
# there needs no change of owner; chown -R changes a symbolic link itself,
|
||||||
|
# never what it points to. exec, so that the signal `sv stop` sends
|
||||||
|
# reaches smallwebwaf itself.
|
||||||
main() {
|
main() {
|
||||||
sleep 1
|
sleep 1
|
||||||
|
chown -R smallwebwaf:smallwebwaf "${SWWAF_STATE_DIR:-/var/lib/smallwebwaf}"
|
||||||
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
|
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user