4 Commits
Author SHA1 Message Date
clawbot df2c5042d2 Keep the bans, the clients and GeoJS's answers in state files (closes #17)
check / check (push) Waiting to run
smallwebwaf now copies its state to bans.json, clients.json and
lookups.json in SWWAF_STATE_DIR, as "Persistent state" in SPEC.md
describes, and reads them back at start, so a restart lifts no ban and
gives no client a fresh allowance. Each client gains a history, and a
ban's notes count the netblock's requests. bans.json is written
SWWAF_STATE_WRITE_DELAY after a ban, and every file every
SWWAF_STATE_COUNTER_INTERVAL and at the stop. A ban read back is masked
to its netblock and refuses every client in it. A file that does not
parse, an unknown version, an entry without a field it needs, or an
unwritable directory stops the start.

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

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

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

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

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

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

Model: opus-5-5
2026-10-06 05:05:31 +02:00
clawbot 50df9ee36e Re-vendor the canonical files from sneak/prompts at dd4027b (closes #65)
check / check (push) Successful in 4m1s
The vendored files are fetched from sneak/prompts commit dd4027b. This
repository's own entries come after the canonical content, at the end of
each file: /bin in .dockerignore, the Go lines of .gitignore and [*.go]
in .editorconfig; the test-support deny list has no entries of its own.
The lint phase moves to golangci-lint v2.14.0. The build stage now takes
the version from git describe on the .git the build context carries,
unless VERSION is passed, and fails when .git is present but no version
comes out. The test phase drops -count=1, which the policy says it does
not need, and keeps its tmpfs build cache. One test calls Header.Get
with X-Real-IP, as canonicalheader asks.

Model: opus-5-5
2026-10-06 04:13:35 +02:00
35 changed files with 4481 additions and 379 deletions
+3 -3
View File
@@ -58,9 +58,6 @@
# Dependencies: restored inside the image, never copied in.
**/node_modules
# The binary `make build` writes on the host; the image builds its own.
/bin
# OS metadata.
**/.DS_Store
**/Thumbs.db
@@ -73,3 +70,6 @@
**/.idea
**/.vscode
**/*.sublime-*
# The binary `make build` writes on the host; the image builds its own.
/bin
+5
View File
@@ -162,6 +162,11 @@ RUN groupadd --system --gid 65532 smallwebwaf \
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
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
# looks too.
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
+220 -107
View File
@@ -13,15 +13,19 @@ JSON log line for every request.
Status: the first two milestones are built
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are the static lists,
which come next in the build order. `smallwebwaf` passes each request to the app
and the app's answer back, unchanged, within its timeouts and size limits, works
out each client's address, refuses a client that sends too many requests, comes
from a country you refuse or from a network you refuse, lets the networks you
choose through, and writes a JSON log line for every request. It comes as the
image the app's own image is built on. The rest of the design comes after that,
in the order of the build order in [`SPEC.md`](SPEC.md). The survey of existing
tools that led to the design is in [`EVALUATION.md`](EVALUATION.md).
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are four parts of
milestone 3: the static lists, the bans that broken rate limits lead to and the
JSON state files, which come next in the build order, and the header size and
the idle time as settings, which come last in it. `smallwebwaf` passes each
request to the app and the app's answer back, unchanged, within its timeouts and
size limits, works out each client's address, bans a client that sends too many
requests, refuses a client that comes from a country you refuse or from a
network you refuse, lets the networks you choose through, keeps its bans, each
client's counters and history, and GeoJS's answers in JSON files across
restarts, and writes a JSON log line for every request. It comes as the image
the app's own image is built on. The rest of the design comes after that, in the
order of the build order in [`SPEC.md`](SPEC.md). The survey of existing tools
that led to the design is in [`EVALUATION.md`](EVALUATION.md).
## Getting started
@@ -43,7 +47,8 @@ works.
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
`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
@@ -58,42 +63,60 @@ and `make run` builds and runs it, listening on port 8080 in front of an app at
is inside, the leftmost is, and with no header the peer is. The app sees what
it would see from traefik directly: the same `Host`, the same
`X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end.
- Enforces the four timeouts and the two size limits below. A limit passed
before the response has started gets `smallwebwaf`'s own answer: `408` for a
client too slow to send its request, `413` for a request body that is too
large, `504` for an app too slow to answer, and `502` for a response that is
too large or an app that cannot be reached. A request that announces a body
over the limit is refused before anything reaches the app. While a request
body is still on its way, a request timeout that runs out answers `408` if
`smallwebwaf` was waiting for the client to send more, and `504` if it was
waiting for the app to take what it had. Once the response has started, a
limit can only cut the connection.
- Enforces the timeouts and the size limits below. A limit passed before the
response has started gets `smallwebwaf`'s own answer: `408` for a client too
slow to send its request, `413` for a request body that is too large, `504`
for an app too slow to answer, and `502` for a response that is too large or
an app that cannot be reached. A request that announces a body over the limit
is refused before anything reaches the app. While a request body is still on
its way, a request timeout that runs out answers `408` if `smallwebwaf` was
waiting for the client to send more, and `504` if it was waiting for the app
to take what it had. Once the response has started, a limit can only cut the
connection.
- Counts each client's requests over a minute, an hour and a day. A request that
takes the client over one of the rate limits below is refused with `429`
before anything reaches the app, and so is each request after it until the
client is back under every limit. A client is one IPv4 address, or one IPv6
/64, since one abuser usually holds a whole /64. Refused requests count too,
so a client that keeps sending too fast stays refused until it slows down.
Each window is counted in two fixed buckets, 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 memory: a restart starts every client
afresh.
- Refuses a request from a country you refuse with `403`, as soon 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 set, each
client's country is looked up through GeoJS (see "Country and AS number
lookup" below); with neither set, no visitor's address leaves the host. A
client on a private, loopback or link-local address has no country and is
takes the client over one of the rate limits below is refused with
`SWWAF_BAN_RESPONSE`, `403` by default, before anything reaches the app, and
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,
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, with their
history, and a restart gives no client a fresh allowance (see "State files"
below).
- Bans a client that breaks a rate limit, as "Bans" in [`SPEC.md`](SPEC.md)
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
81 hours; a ban that would last longer than seven days is permanent instead. A
ban covers the client's netblock: its IPv4 address, or the netblock around it
that `SWWAF_BAN_SCOPE_V4_PREFIX` sets, or its IPv6 /64. While it lasts, every
request from the netblock is refused with `SWWAF_BAN_RESPONSE` after the
static lists and before the country lists, so the client is not looked up, and
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
window and the requests counted in it, the request that broke it, the client's
country when it was looked up, the netblock's requests since it was first
seen, how many of them the ban has refused, and how many bans the netblock had
before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent;
past that, the earliest ban of the netblock that has gone longest without a
request is dropped first. `bans.json` shows the bans and their notes, and a
restart lifts none (see "State files" below); lifting a ban by editing it
comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
- 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
is not counted for the rate limits. While one of the country lists below is
set, each client's country is looked up through GeoJS (see "Country and AS
number lookup" below); with neither set, no visitor's address leaves the host.
A client on a private, loopback or link-local address has no country and is
never looked up: `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses it unless it is
in `SWWAF_ALLOW_NETS`, and `SWWAF_DENIED_COUNTRIES` does not refuse it.
- Checks the client's own address against the static lists, the three netblock
settings below, before anything else, its country included. A client in
`SWWAF_ALLOW_NETS` skips the country lists and the rate limits, and is not
looked up; the timeouts and size limits still apply. A client in
`SWWAF_DENY_NETS` is refused with `403` before its body is read, and the
request is not counted for the rate limits; an address in `SWWAF_ALLOW_NETS`
too is let through. A client in `SWWAF_RATE_LIMIT_EXEMPT_NETS` is neither
counted nor refused by the rate limits; the country lists still apply to it.
`SWWAF_ALLOW_NETS` skips bans, the country lists and the rate limits, and is
not looked up; the timeouts and size limits still apply. A client in
`SWWAF_DENY_NETS` is refused with `SWWAF_BAN_RESPONSE` before its body is
read, and the request is not counted for the rate limits; an address in
`SWWAF_ALLOW_NETS` too is let through. A client in
`SWWAF_RATE_LIMIT_EXEMPT_NETS` is neither counted nor refused by the rate
limits; the country lists and bans still apply to it.
- Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any
check and without asking the app, for the image's health check.
- Writes a line in the request log for each request (see "Request log" below).
@@ -113,6 +136,16 @@ it, and the effective settings are logged at start.
- `SWWAF_CLIENT_REQUEST_TIMEOUT` (default `60s`): how long a client may take to
send its request line and headers, and then, from the end of the headers, its
body.
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest request
line and headers a client may send. Over it, the answer is `431` and nothing
reaches the app. It must be more than `4K`, and cannot be `off`: Go's HTTP
server always has such a limit, and reads 4 KiB past the one it is given
before it refuses.
- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open connection
may wait for its next request before `smallwebwaf` closes it. The default is
longer than the 90 seconds after which traefik closes a connection it is not
using, so traefik never sends a request on a connection `smallwebwaf` is
closing.
- `SWWAF_CLIENT_RESPONSE_TIMEOUT` (default `30m`): how long the response may
take to reach the client, from the end of the request to the last byte.
- `SWWAF_UPSTREAM_REQUEST_TIMEOUT` (default `60s`): how long connecting to the
@@ -121,8 +154,9 @@ it, and the effective settings are logged at start.
to send its whole answer, from the end of the request to the last byte.
- `SWWAF_REQUEST_MAX_BYTES` (default `100M`): the largest request body.
- `SWWAF_RESPONSE_MAX_BYTES` (default `5G`): the largest response body.
- `SWWAF_ALLOW_NETS` (default empty): netblocks whose clients skip the country
lists and the rate limits, such as your monitoring or your own networks.
- `SWWAF_ALLOW_NETS` (default empty): netblocks whose clients skip bans, the
country lists and the rate limits, such as your monitoring or your own
networks.
- `SWWAF_RATE_LIMIT_EXEMPT_NETS` (default empty): netblocks whose clients the
rate limits do not apply to, such as a machine that talks to the app all day.
- `SWWAF_DENY_NETS` (default empty): netblocks whose clients are always refused.
@@ -137,6 +171,29 @@ it, and the effective settings are logged at start.
countries whose clients get through, for example `us,de`. A client whose
country cannot be found is refused too, so that new clients are not let in
whenever GeoJS stops answering.
- `SWWAF_BAN_RESPONSE` (default `403`): how a refused client is answered, one
that is banned, breaks a rate limit, is in `SWWAF_DENY_NETS` or comes from a
refused country: `403`, `429`, or `close` to close the connection without an
answer. Behind traefik, `close` does not leave the client unanswered: traefik
answers `502`, as it does whenever its backend drops a connection.
- `SWWAF_LIMIT_BAN_DURATION` (default `1h`): the ban for a first broken rate
limit.
- `SWWAF_LIMIT_BAN_REPEAT_WINDOW` (default `24h`): a rate limit broken again
within this time after a ban ended bans for three times as long as that ban.
- `SWWAF_MAX_BAN_DURATION` (default `7d`): a ban that would be longer is
permanent instead.
- `SWWAF_MAX_BANS` (default `5000`): the most bans kept, past, active and
permanent.
- `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
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
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
@@ -145,16 +202,14 @@ and a bare address stands for itself alone. Countries are the two-letter codes
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
`su`, stops the start, and so does a code on both country lists. `off` switches
a timeout, a size limit or a rate limit off.
a timeout, a size limit or a rate limit off;
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings and the state settings
cannot be off.
Several limits are fixed rather than settings. The request line and headers may
take up to 32 KiB, above which the answer is `431` and nothing reaches the app.
A kept-open connection that sends nothing for 120 seconds is closed. That is
longer than the 90 seconds after which traefik closes a connection it is not
using, so traefik never sends a request on a connection `smallwebwaf` is
closing. At most 20,000 clients are kept for the rate limits, and an IPv6 client
is counted by its /64. A new client waits at most a second for its country, and
at most 100,000 answers from GeoJS are kept, for 7 days each.
Several limits are fixed rather than settings. At most 20,000 clients are kept,
with their counters and history, and an IPv6 client is counted by its /64. A new
client waits at most a second for its country, and at most 100,000 answers from
GeoJS are kept, for 7 days each.
## Request log
@@ -167,23 +222,27 @@ refused ones included:
- `time` is when the request arrived, in UTC. `peer_ip` is the TCP peer,
normally traefik. `path` and `query` are as the client sent them.
- `country` is the client's country as GeoJS places it, and empty when it is not
known: with neither country list set, for a client in `SWWAF_ALLOW_NETS` or
`SWWAF_DENY_NETS`, for a client on a private, loopback or link-local address,
and when GeoJS cannot place the client or has not answered in time.
- `country` is the client's country as GeoJS places it. It is empty with neither
country list set, for a client in `SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS`, for
a client on a private, loopback or link-local address, when GeoJS cannot place
the client or has not answered in time, and for a request refused because a
ban covers its client, even when the client's country is known.
- `status` is what the client was sent, `0` if nothing was; `upstream_status` is
what the app answered, and is left out when the app did not answer.
- `request_bytes` and `response_bytes` count body bytes.
- `action` is `forward` for a request passed to the app, `denied` for one
refused because its client is in `SWWAF_DENY_NETS`, `country_denied` for one
refused for its client's country, `rate_limited` for one refused for a rate
limit, `too_large` for a request or response over its size limit, `timed_out`
for one that ran out of time, `upstream_error` when the app could not be
reached or its answer broke off, and `admin` for one `smallwebwaf` answered at
its own endpoint.
- `limit_hit` is there for a request refused for a rate limit, and names the
refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused
because a ban covers its client, `country_denied` for one refused for its
client's country, `rate_limited` for one that broke a rate limit and banned
its client, `too_large` for a request or response over its size limit,
`timed_out` for one that ran out of time, `upstream_error` when the app could
not be reached or its answer broke off, and `admin` for one `smallwebwaf`
answered at its own endpoint.
- `limit_hit` is there for a request that broke a rate limit, and names the
window whose limit it went over: `minute`, `hour` or `day`, the shortest if it
went over several.
went over several. `offence` is then `limit`.
- `ban_expires` is there for a request that made a ban or was refused under one,
and gives when the ban ends, in the same form as `time`, or `permanent`.
- `aborted` is there, and true, when the client went away early.
- `duration_total` and `duration_upstream_total` are in milliseconds.
@@ -193,10 +252,52 @@ settings, stop, errors) share the stream as JSON lines marked
Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and
headers before `smallwebwaf` sees the request, and some requests end there,
without a line in the log: headers over 32 KiB, which it answers `431`, headers
slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`, whose connection it closes without
an answer, and requests it cannot read at all, which it answers itself, mostly
with `400`.
without a line in the log: headers over `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`,
which it answers `431`, headers slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`,
whose connection it closes without an answer, and requests it cannot read at
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
@@ -299,9 +400,9 @@ goes through the candidates one by one.
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
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));
until then the rate counters and the GeoJS answers are kept in memory only,
and a restart loses them.
for the bans, the clients and the GeoJS answers are built (see "State files"
above); the others come with their features, and taking in an edit while
running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
- 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
`/_smallwebwaf/` on the app's own address, through traefik like any other
@@ -397,8 +498,9 @@ main "$@"
the health check on `127.0.0.1`.
- `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;
without one, it still starts. The state files come in milestone 3 or later;
until then it writes nothing to disk and needs no volume.
without one, it still starts. At each start the `run` script of `smallwebwaf`
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
requests and gives those in progress five seconds to finish.
@@ -418,28 +520,28 @@ the metrics, failure behaviour and the build order.
So far `smallwebwaf` looks up only the country, only through GeoJS, and only
while `SWWAF_DENIED_COUNTRIES` or `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` is set:
then the address of every new visitor outside `SWWAF_ALLOW_NETS` and
`SWWAF_DENY_NETS` is sent to GeoJS, and with neither set, none is. An IPv6
visitor is asked about by the first address of its /64. A new visitor waits at
most a second for its answer, and without one counts as coming from an unknown
country until the answer arrives. The addresses waiting are asked about
together, up to 200 in one request, one request at a time; at most 10,000
visitors wait, and one more counts as coming from an unknown country until there
is room. While GeoJS fails, visitors with a kept answer are unaffected and new
ones count as coming from an unknown country. GeoJS is then left alone for a
second, twice as long after each further failure up to five minutes, and asked
again by the next request that needs it.
then the address of every new visitor is sent to GeoJS, except a visitor in
`SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS` and one refused because a ban covers its
netblock, and with neither set, none is. An IPv6 visitor is asked about by the
first address of its /64. A new visitor waits at most a second for its answer,
and without one counts as coming from an unknown country until the answer
arrives. The addresses waiting are asked about together, up to 200 in one
request, one request at a time; at most 10,000 visitors wait, and one more
counts as coming from an unknown country until there is room. While GeoJS fails,
visitors with a kept answer are unaffected and new ones count as coming from an
unknown country. GeoJS is then left alone for a second, twice as long after each
further failure up to five minutes, and asked again by the next request that
needs it.
In the full design, `smallwebwaf` looks up the AS number and country of every
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
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
is kept in memory for seven days, and many addresses are asked about in one
request; writing the answers to disk, so that they survive a restart, comes in
milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md)). GeoJS
publishes no rate limit but may block a caller it thinks asks too much; while it
is not answering, new visitors count as coming from an unknown country, which
is kept for seven days, in memory and in `lookups.json`, so that it survives a
restart, and many addresses are asked about in one request. GeoJS publishes no
rate limit but may block a caller it thinks asks too much; while it is not
answering, new visitors count as coming from an unknown country, which
`SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses.
To keep your visitors' addresses on your own host, set
@@ -470,21 +572,26 @@ addresses are never sent to GeoJS.
## How the code is laid out
- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`.
- `internal/smallwebwaf`: the process: it reads the settings, listens, serves
requests until `SIGTERM` or `SIGINT`, and stops. Run as
`smallwebwaf healthcheck`, it is the image's health check instead.
- `internal/smallwebwaf`: the process: it reads the settings and the state
files, listens, serves requests until `SIGTERM` or `SIGINT`, and stops,
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/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
standard library's `httputil.ReverseProxy` within the timeouts and size
limits, and writes the request's log line. Its `check` method is where a
request is refused before anything reaches the app: for `SWWAF_DENY_NETS`, for
the country lists, for a rate limit, and for an announced body over the size
limit.
a ban, for the country lists, for a rate limit, which bans the client, and for
an announced body over the size limit.
- `internal/bans`: the ban ledger: each netblock's bans with their notes, how
long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held.
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the
answers.
- `internal/ratelimit`: counts each client's requests and tells when one takes
it over a rate limit.
- `internal/ratelimit`: the table of clients: counts each client's requests,
tells when one takes it over a rate limit, and keeps each client's history.
- `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
process's own messages.
- `Dockerfile`: the lint and test phases, then the image, whose last stage
@@ -494,8 +601,9 @@ addresses are never sent to GeoJS.
checks.
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
table of clients to 20,000 and the GeoJS answers to 100,000, dropping the least
recently seen. The country codes are the list in `internal/config/config.go`.
table of clients to 20,000, the GeoJS answers to 100,000 and the banned
netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen. The country
codes are the list in `internal/config/config.go`.
## Entrypoints
@@ -524,19 +632,24 @@ so that they run in minimal containers.
- `script/install-precommit`: installs that hook; `make hooks` runs it.
- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for
working on the code by hand; `make build` runs it.
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it;
`make run` runs it.
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it, with
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
`deploy/example-app`, runs it, and checks that the health check passes, that a
request reaches the app through `smallwebwaf`, and that `sv stop` and
`docker stop` stop it in order; then removes the container and both images. It
needs network access, for nixpkgs' binary cache, and `script/check` does not
run it; `make example-app` does.
`deploy/example-app`, runs it with a volume for the state files, and checks
that the health check passes, that a request reaches the app through
`smallwebwaf`, that a second request in a minute bans the client, that
`sv stop` and `docker stop` stop it in order, and that a new container on the
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
- The rest of milestone 3, after the static lists, and the rest of the design,
in the order of the build order in [`SPEC.md`](SPEC.md).
- The rest of milestone 3, from taking in an admin's edits to the state files
(https://git.eeqj.de/sneak/smallwebwaf/issues/68) up to the metrics endpoint,
and the rest of the design, in the order of the build order in
[`SPEC.md`](SPEC.md).
## Documents
+15 -12
View File
@@ -293,7 +293,8 @@ it.
needs: an alert destination, an account key, a token.
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
container, and so its environment variables, with the app it protects.
- Any limit or threshold can be switched off with the value `off`.
- Any limit or threshold can be switched off with the value `off`, except
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`.
- A list set to an empty value is an empty list, and replaces the default.
- Every setting may instead be given as a file holding the value, named by the
setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for
@@ -413,7 +414,9 @@ The settings, by group:
headers, its body.
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
request line and headers a client may send. Over it, `smallwebwaf` answers
`431` and closes the connection, and nothing reaches the app.
`431` and closes the connection, and nothing reaches the app. It must be
more than `4K`, and cannot be `off`: Go's HTTP server always has such a
limit, and reads 4 KiB past the one it is given before it refuses.
- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open
connection may wait for its next request before `smallwebwaf` closes it.
It is longer than the 90 seconds after which traefik, by default, closes a
@@ -945,9 +948,9 @@ and the running `smallwebwaf` takes the edit in.
- what was broken: the rule ids and target that matched, or the limit, its
window, the count reached and the client's limit percentage with what set
it; and any reputation sources that listed the client;
- the requests that caused the ban, up to the last ten: time, method, host,
path with its query string, status and user agent, each text cut to 256
bytes;
- the request that caused the ban, the one that broke the limit or carried
the clear sign of attack: time, method, host, path with its query string,
status and user agent, each text cut to 256 bytes;
- how many requests counted toward the ban, and the time span over which
they came;
- the netblock's total requests since it was first seen, and the requests
@@ -958,13 +961,13 @@ and the running `smallwebwaf` takes the edit in.
the table is full, so on a public service the file grows to the default
`SWWAF_MAX_TRACKED_CLIENTS` of 20,000, about 20 MiB. Written every 15
minutes, that is under 2 GiB of disk writes a day.
- `bans.json` takes about 2 KiB per ban and at most about 8 KiB, since the
texts in the notes are cut short. At the default `SWWAF_MAX_BANS` of 5,000
it is about 10 MiB, and never more than about 40 MiB, plus whatever bans
an admin made. It is written when a ban is made, lifted or made permanent,
at most once every 10 seconds, and otherwise with the 15-minute write, so
its writes follow the bans made: with a full file, a hundred new bans a
day come to about 1 GiB of disk writes.
- `bans.json` takes about 1.2 KiB per ban and at most about 2.5 KiB, since
the notes hold one request and their texts are cut short. At the default
`SWWAF_MAX_BANS` of 5,000 it is about 6 MiB, and never more than about 12
MiB, plus whatever bans an admin made. It is written when a ban is made,
lifted or made permanent, at most once every 10 seconds, and otherwise
with the 15-minute write, so its writes follow the bans made: with a full
file, a hundred new bans a day come to about 600 MiB of disk writes.
- `lookups.json` takes about 150 bytes per answer, about 15 MiB when full.
Written every 15 minutes, that is under 1.5 GiB of disk writes a day.
- `reputation.json` and `alerts.json` are usually a few MiB or less.
+351
View File
@@ -0,0 +1,351 @@
// 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
// "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
import (
"net/netip"
"slices"
"strings"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
)
// repeatFactor is how many times as long as the netblock's last ban a ban
// for a limit broken again within the repeat window lasts.
const repeatFactor = 3
// maxTextBytes is how much of each text in a ban's notes is kept.
const maxTextBytes = 256
// Rules are how long a ban for a broken limit lasts, and how many bans
// are held.
type Rules struct {
// LimitBanDuration is how long a first ban lasts.
LimitBanDuration time.Duration
// LimitBanRepeatWindow is how soon after the netblock's last ban
// ended a broken limit counts as a repeat, which bans for
// repeatFactor times as long as that ban.
LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban; a ban that would be longer is
// permanent instead.
MaxBanDuration time.Duration
// MaxBans is the most bans held, at least one. Past it, the earliest
// ban of the netblock that has gone longest without a request is
// dropped.
MaxBans int
}
// Ban is a ban on a netblock for a broken limit, the only kind of ban
// smallwebwaf makes so far.
type Ban struct {
Netblock netip.Prefix
Start time.Time
// Expires is when the ban ends, zero for a permanent ban.
Expires time.Time
Notes Notes
}
// Permanent reports whether the ban never runs out.
func (b Ban) Permanent() bool {
return b.Expires.IsZero()
}
// ActiveAt reports whether the ban refuses requests at now.
func (b Ban) ActiveAt(now time.Time) bool {
return b.Permanent() || now.Before(b.Expires)
}
// 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 {
// Country is the client's country, when it was looked up.
Country string `json:"country"`
// Limit, Window and Count are the limit that was broken, its window,
// "minute", "hour" or "day", and the count reached: the client's
// requests in the window, the one that broke the limit included.
// These are the requests that counted toward the ban, and the window
// is the time over which they came.
Limit int64 `json:"limit"`
Window string `json:"window"`
Count float64 `json:"count"`
// Request is the request that broke the limit.
Request Request `json:"request"`
// Requests is how many requests the netblock has sent since it was
// first seen, and Refused how many of them the ban has refused so
// far. Both go up with each request the ban refuses.
Requests int64 `json:"requests"`
Refused int64 `json:"refused"`
// EarlierBans is how many bans the netblock had before this one.
EarlierBans int `json:"earlier_bans"`
}
// 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 {
Time time.Time `json:"time"`
Method string `json:"method"`
Host string `json:"host"`
// Path is the path with its query string.
Path string `json:"path"`
// Status is what the client was sent, 0 if nothing was.
Status int `json:"status"`
UserAgent string `json:"user_agent"`
}
// Ledger holds the bans. It is safe for concurrent use.
type Ledger struct {
rules Rules
// changed receives a value when a ban is made, unless one is waiting
// already.
changed chan struct{}
mu sync.Mutex
// netblocks holds each banned netblock's bans, oldest first. Check
// makes each netblock it finds the most recently seen.
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans.
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.
func New(rules Rules) *Ledger {
// Every netblock held has a ban, so there are never more netblocks
// than rules.MaxBans, and the LRU never drops one itself.
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](rules.MaxBans, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
return &Ledger{
rules: rules,
changed: make(chan struct{}, 1),
netblocks: netblocks,
}
}
// Changed receives a value after a ban is made, so that bans.json can be
// written. Several bans made before it is read leave one value.
func (l *Ledger) Changed() <-chan struct{} {
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()
defer l.mu.Unlock()
lengths := l.v6Lengths
if client.Is4() {
lengths = l.v4Lengths
}
for _, length := range lengths {
bans, found := l.netblocks.Get(netip.PrefixFrom(client, length).Masked())
if !found {
continue
}
// 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
}
}
return Ban{}, false
}
// BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban. A first ban lasts LimitBanDuration. A ban made within
// LimitBanRepeatWindow after the netblock's last ban ended lasts
// repeatFactor times as long as that one. A ban that would be longer
// than MaxBanDuration is permanent instead. If a ban on netblock is still
// active, as when two of its requests break a limit at once, that ban is
// returned and no other is made. The ledger fills in the notes' Refused
// and EarlierBans itself.
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
l.mu.Lock()
defer l.mu.Unlock()
var last *Ban
bans, found := l.netblocks.Get(netblock)
if found {
last = &(*bans)[len(*bans)-1]
if last.ActiveAt(now) {
return *last
}
notes.EarlierBans = last.Notes.EarlierBans + 1
}
notes.Request = notes.Request.cut()
ban := Ban{
Netblock: netblock,
Start: now,
Expires: l.expiry(last, now),
Notes: notes,
}
l.add(ban)
select {
case l.changed <- struct{}{}:
default: // a value is waiting already
}
return ban
}
// Bans returns the bans held on netblock, oldest first. It is not a
// request from netblock, and leaves when it was last seen unchanged.
func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
bans, found := l.netblocks.Peek(netblock)
if !found {
return nil
}
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
// when it is permanent. last is the netblock's last ban, which has ended,
// or nil when it has none.
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration
if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow {
lastLength := last.Expires.Sub(last.Start)
// This is repeatFactor * lastLength > MaxBanDuration, written so
// that it cannot overflow.
if lastLength > l.rules.MaxBanDuration/repeatFactor {
return time.Time{}
}
length = repeatFactor * lastLength
}
if length > l.rules.MaxBanDuration {
return time.Time{}
}
return now.Add(length)
}
// dropOne drops the earliest ban of the netblock that has gone longest
// without a request, and the netblock with it if that was its only ban.
func (l *Ledger) dropOne() {
netblock, bans, _ := l.netblocks.GetOldest()
if len(*bans) == 1 {
l.netblocks.Remove(netblock)
} else {
*bans = slices.Delete(*bans, 0, 1)
}
l.held--
}
// cut returns r with each text cut to maxTextBytes and copied, so that
// the notes do not keep the rest of the request in memory.
func (r Request) cut() Request {
r.Method = cutText(r.Method)
r.Host = cutText(r.Host)
r.Path = cutText(r.Path)
r.UserAgent = cutText(r.UserAgent)
return r
}
// cutText returns a copy of the first maxTextBytes of text.
func cutText(text string) string {
return strings.Clone(text[:min(len(text), maxTextBytes)])
}
+269
View File
@@ -0,0 +1,269 @@
package bans_test
import (
"net/netip"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
const day = 24 * time.Hour
func TestRepeatsTripleUntilPermanent(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
// 81 hours.
for i, hours := range []int{1, 3, 9, 27, 81} {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
length := time.Duration(hours) * time.Hour
if !ban.Expires.Equal(now.Add(length)) || ban.Notes.EarlierBans != i {
t.Fatalf("ban %d lasts %s with %d earlier bans, want %d hours and %d",
i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i)
}
now = ban.Expires
}
// The sixth would last 243 hours, more than seven days: it is
// permanent, and never ends.
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
}
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day))
if !banned {
t.Error("a permanent ban ended")
}
}
func TestRepeatWindowRunsOut(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// gap is the time between the end of the first ban and the second.
gap time.Duration
want time.Duration
}{
{"broken again as the window ends", day, 3 * time.Hour},
{"broken again after the window", day + time.Nanosecond, time.Hour},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
if second.Expires.Sub(second.Start) != tc.want || second.Notes.EarlierBans != 1 {
t.Errorf("second ban lasts %s with %d earlier bans, want %s and 1",
second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want)
}
})
}
}
func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
ledger := bans.New(rules)
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
if !ban.Permanent() {
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
}
}
func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
t.Parallel()
// With bans of up to 100,000 days, the 14th ban in a row, of 3^13
// hours, is within the maximum, and three times as long would not fit
// in a time.Duration. The 15th is permanent.
rules := defaultRules()
rules.MaxBanDuration = 100000 * day
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
for i := range 14 {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Expires.After(ban.Start) {
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
}
now = ban.Expires
}
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
}
}
func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
if again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
again, len(ledger.Bans(netblock)), first)
}
}
func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
for range 3 {
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got.Start != ban.Start {
t.Fatalf("check during the ban gives %+v and %t", got, banned)
}
}
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
if banned {
t.Error("another netblock is banned")
}
_, banned = ledger.Check(netblock.Addr(), ban.Expires)
if banned {
t.Error("the ban did not end")
}
// The netblock's requests went from 5 to 8 with the three refused.
notes := ledger.Bans(netblock)[0].Notes
if notes.Refused != 3 || notes.Requests != 8 {
t.Errorf("the notes count %d refused requests of %d, want 3 of 8",
notes.Refused, notes.Requests)
}
}
func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.MaxBans = 3
ledger := bans.New(rules)
a := netip.MustParsePrefix("203.0.113.1/32")
b := netip.MustParsePrefix("203.0.113.2/32")
c := netip.MustParsePrefix("203.0.113.3/32")
d := netip.MustParsePrefix("2001:db8::/64")
now := midnight()
first := ledger.BanForLimit(a, now, bans.Notes{})
ledger.BanForLimit(b, now, bans.Notes{})
ledger.BanForLimit(c, now, bans.Notes{})
// A request from a makes b the netblock seen longest ago, and its ban
// goes to make room for d's.
ledger.Check(a.Addr(), now)
ledger.BanForLimit(d, now, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
// a is banned again once its ban has ended; c, seen longest ago, goes.
ledger.BanForLimit(a, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 2, c: 0, d: 1})
// With d seen since, a is seen longest ago, and its earlier ban goes
// first.
ledger.Check(d.Addr(), first.Expires)
ledger.BanForLimit(b, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
if !ledger.Bans(a)[0].Start.Equal(first.Expires) {
t.Errorf("a kept its ban of %s, want the later one", ledger.Bans(a)[0].Start)
}
}
func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
t.Parallel()
// With room for one ban, the netblock's ended ban goes to make room for
// its new one, whose notes still count it.
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
held := ledger.Bans(netblock)
if len(held) != 1 || held[0] != second || held[0].Notes.EarlierBans != 1 {
t.Errorf("the ledger holds %+v, want only the second ban, with 1 earlier ban",
held)
}
}
func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
long := strings.Repeat("a", 300)
request := bans.Request{
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
}
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
cut := long[:256]
want := bans.Request{
Time: midnight(), Method: cut, Host: cut, Path: cut, Status: 403, UserAgent: cut,
}
if ban.Notes.Request != want || ledger.Bans(netblock)[0].Notes.Request != want {
t.Errorf("the notes keep %+v, want each text cut to 256 bytes", ban.Notes.Request)
}
}
// defaultRules are the rules at the settings' defaults.
func defaultRules() bans.Rules {
return bans.Rules{
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: day,
MaxBanDuration: 7 * day,
MaxBans: 5000,
}
}
// midnight is when the tests' first bans are made.
func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
}
// wantBans checks how many bans the ledger holds on each netblock.
func wantBans(t *testing.T, ledger *bans.Ledger, want map[netip.Prefix]int) {
t.Helper()
for netblock, count := range want {
got := len(ledger.Bans(netblock))
if got != count {
t.Errorf("%s has %d bans, want %d", netblock, got, count)
}
}
}
+193
View File
@@ -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)
}
}
+176 -5
View File
@@ -9,8 +9,10 @@ import (
"log/slog"
"math"
"net"
"net/http"
"net/netip"
"net/url"
"path/filepath"
"slices"
"strconv"
"strings"
@@ -30,6 +32,13 @@ type Config struct {
// ClientRequestTimeout bounds reading the whole request from the
// client (SWWAF_CLIENT_REQUEST_TIMEOUT).
ClientRequestTimeout time.Duration
// ClientRequestHeaderMaxBytes is the largest request line and headers
// a client may send (SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES). It is
// never off, and always more than 4K.
ClientRequestHeaderMaxBytes int64
// ClientIdleTimeout bounds how long a kept-open client connection
// may wait for its next request (SWWAF_CLIENT_IDLE_TIMEOUT).
ClientIdleTimeout time.Duration
// ClientResponseTimeout bounds writing the whole response to the
// client (SWWAF_CLIENT_RESPONSE_TIMEOUT).
ClientResponseTimeout time.Duration
@@ -67,6 +76,33 @@ type Config struct {
// capitals, as GeoJS gives them.
DeniedCountries []string
ExclusivelyAllowedCountries []string
// BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
// breaks a rate limit, SWWAF_DENY_NETS and the country lists.
BanResponse int
// LimitBanDuration is the ban for a first broken rate limit
// (SWWAF_LIMIT_BAN_DURATION). A limit broken again within
// LimitBanRepeatWindow after the last ban ended bans for three times
// as long as that ban (SWWAF_LIMIT_BAN_REPEAT_WINDOW), and a ban that
// would be longer than MaxBanDuration is permanent instead
// (SWWAF_MAX_BAN_DURATION). None of them can be off.
LimitBanDuration time.Duration
LimitBanRepeatWindow time.Duration
MaxBanDuration time.Duration
// MaxBans is the most bans held (SWWAF_MAX_BANS).
MaxBans int
// BanScopeV4Prefix is the length of the netblock around an IPv4
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
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
// log line at start.
@@ -82,6 +118,7 @@ const (
kibibyte = 1 << 10
mebibyte = 1 << 20
gibibyte = 1 << 30
ipv4Bits = 32
)
var (
@@ -102,7 +139,17 @@ var (
"such as http://127.0.0.1:8081")
errNotCountry = errors.New(
"is not a two-letter country code such as de or kp")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New(
"is not a duration above zero, such as 1h or 7d")
errNotNumberAboveZero = errors.New(
"is not a whole number above zero, such as 5000")
errNotBanResponse = errors.New("is not 403, 429 or close")
errNotV4Prefix = errors.New(
"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
@@ -111,10 +158,13 @@ var (
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv}
cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
ClientRequestHeaderMaxBytes: env.headerSize(
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES", "32K"),
ClientIdleTimeout: env.duration("SWWAF_CLIENT_IDLE_TIMEOUT", "120s"),
ClientResponseTimeout: env.duration("SWWAF_CLIENT_RESPONSE_TIMEOUT", "30m"),
UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"),
UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"),
@@ -129,6 +179,15 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
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 {
@@ -225,6 +284,15 @@ func (e *environment) size(name, defaultValue string) int64 {
return size
}
// headerSize reads the setting that is the largest request line and
// headers.
func (e *environment) headerSize(name, defaultValue string) int64 {
size, err := parseHeaderSize(e.value(name, defaultValue))
e.check(name, err)
return size
}
// count reads a setting that is a number of requests.
func (e *environment) count(name, defaultValue string) int64 {
count, err := parseCount(e.value(name, defaultValue))
@@ -241,6 +309,50 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries
}
// durationNotOff reads a setting that is a duration and, unlike a
// timeout, cannot be off.
func (e *environment) durationNotOff(name, defaultValue string) time.Duration {
duration, err := parseDurationNotOff(e.value(name, defaultValue))
e.check(name, err)
return duration
}
// numberNotOff reads a setting that is a whole number above zero, which
// cannot be off.
func (e *environment) numberNotOff(name, defaultValue string) int {
number, err := parseNumberNotOff(e.value(name, defaultValue))
e.check(name, err)
return number
}
// banResponse reads a setting that is how a refused client is answered.
func (e *environment) banResponse(name, defaultValue string) int {
status, err := parseBanResponse(e.value(name, defaultValue))
e.check(name, err)
return status
}
// v4Prefix reads a setting that is the length of an IPv4 netblock.
func (e *environment) v4Prefix(name, defaultValue string) int {
length, err := parseV4Prefix(e.value(name, defaultValue))
e.check(name, err)
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
// whole number of days such as 7d, or off.
func parseDuration(value string) (time.Duration, error) {
@@ -296,6 +408,19 @@ func parseSize(value string) (int64, error) {
return n * unit, nil
}
// parseHeaderSize reads the largest request line and headers: a size as
// parseSize reads it, but more than 4K and never off. Go's server reads 4K
// past the limit it is given before it refuses, so proxy.New gives it this
// size less 4K, which must leave a limit.
func parseHeaderSize(value string) (int64, error) {
size, err := parseSize(value)
if err != nil || size <= 4*kibibyte {
return 0, fmt.Errorf("%q %w", value, errNotOver4K)
}
return size, nil
}
// splitUnit splits a size into its number and the bytes its suffix
// stands for.
func splitUnit(value string) (string, int64) {
@@ -329,6 +454,52 @@ func parseCount(value string) (int64, error) {
return n, nil
}
// parseDurationNotOff reads a duration above zero, as parseDuration does,
// but not off.
func parseDurationNotOff(value string) (time.Duration, error) {
duration, err := parseDuration(value)
if err != nil || duration == 0 {
return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero)
}
return duration, nil
}
// parseNumberNotOff reads a whole number above zero.
func parseNumberNotOff(value string) (int, error) {
n, err := strconv.Atoi(value)
if err != nil || n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotNumberAboveZero)
}
return n, nil
}
// parseBanResponse reads how a refused client is answered: 403, 429, or
// close, which is 0.
func parseBanResponse(value string) (int, error) {
switch value {
case "403":
return http.StatusForbidden, nil
case "429":
return http.StatusTooManyRequests, nil
case "close":
return 0, nil
default:
return 0, fmt.Errorf("%q %w", value, errNotBanResponse)
}
}
// parseV4Prefix reads the length of an IPv4 netblock, from 0 to 32.
func parseV4Prefix(value string) (int, error) {
n, err := strconv.Atoi(value)
if err != nil || n < 0 || n > ipv4Bits {
return 0, fmt.Errorf("%q %w", value, errNotV4Prefix)
}
return n, nil
}
// parseList splits a comma-separated list and trims the spaces around
// each item. An empty value is an empty list.
func parseList(value string) ([]string, error) {
+152 -23
View File
@@ -20,6 +20,8 @@ const (
upstreamURL = "SWWAF_UPSTREAM_URL"
trustedProxies = "SWWAF_TRUSTED_PROXIES"
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
@@ -33,6 +35,15 @@ const (
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
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.
@@ -66,16 +77,27 @@ func TestDefaults(t *testing.T) {
cfg := fromEnvironment(t, environment{})
wantSettings(t, cfg, config.Config{
ListenAddr: ":8080",
ClientRequestTimeout: time.Minute,
ClientResponseTimeout: 30 * time.Minute,
UpstreamRequestTimeout: time.Minute,
UpstreamResponseTimeout: 30 * time.Minute,
RequestMaxBytes: 100 << 20,
ResponseMaxBytes: 5 << 30,
RateLimitPerMinute: 1000,
RateLimitPerHour: 10000,
RateLimitPerDay: 50000,
ListenAddr: ":8080",
ClientRequestTimeout: time.Minute,
ClientRequestHeaderMaxBytes: 32 << 10,
ClientIdleTimeout: 2 * time.Minute,
ClientResponseTimeout: 30 * time.Minute,
UpstreamRequestTimeout: time.Minute,
UpstreamResponseTimeout: 30 * time.Minute,
RequestMaxBytes: 100 << 20,
ResponseMaxBytes: 5 << 30,
RateLimitPerMinute: 1000,
RateLimitPerHour: 10000,
RateLimitPerDay: 50000,
BanResponse: 403,
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute,
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -99,6 +121,8 @@ func TestValuesAsSet(t *testing.T) {
upstreamURL: "https://app.internal:8443/",
trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32",
clientRequestTimeout: "90s",
clientHeaderMaxBytes: "8K",
clientIdleTimeout: "5m",
clientResponseTimeout: "7d",
upstreamRequestTimeout: "1h30m",
upstreamResponseTimeout: off,
@@ -112,19 +136,39 @@ func TestValuesAsSet(t *testing.T) {
rateLimitPerDay: "6000",
deniedCountries: "cn, RU,kp,Xk",
allowedCountries: "de",
banResponse: "429",
limitBanDuration: "15m",
limitBanRepeatWindow: "2d",
maxBanDuration: "30d",
maxBans: "100",
banScopeV4Prefix: "24",
stateDir: "/srv/waf-state",
stateWriteDelay: "500ms",
stateCounterInterval: "1h",
})
wantSettings(t, cfg, config.Config{
ListenAddr: "127.0.0.1:9000",
ClientRequestTimeout: 90 * time.Second,
ClientResponseTimeout: 7 * 24 * time.Hour,
UpstreamRequestTimeout: 90 * time.Minute,
UpstreamResponseTimeout: 0,
RequestMaxBytes: 512 << 10,
ResponseMaxBytes: 1234,
RateLimitPerMinute: 60,
RateLimitPerHour: 600,
RateLimitPerDay: 6000,
ListenAddr: "127.0.0.1:9000",
ClientRequestTimeout: 90 * time.Second,
ClientRequestHeaderMaxBytes: 8 << 10,
ClientIdleTimeout: 5 * time.Minute,
ClientResponseTimeout: 7 * 24 * time.Hour,
UpstreamRequestTimeout: 90 * time.Minute,
UpstreamResponseTimeout: 0,
RequestMaxBytes: 512 << 10,
ResponseMaxBytes: 1234,
RateLimitPerMinute: 60,
RateLimitPerHour: 600,
RateLimitPerDay: 6000,
BanResponse: 429,
LimitBanDuration: 15 * time.Minute,
LimitBanRepeatWindow: 48 * time.Hour,
MaxBanDuration: 30 * 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour,
})
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -163,12 +207,42 @@ func TestSizesAndOff(t *testing.T) {
requestMaxBytes: "3G",
responseMaxBytes: off,
clientRequestTimeout: off,
clientIdleTimeout: off,
})
if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 ||
cfg.ClientRequestTimeout != 0 {
t.Errorf("3G, off and off read as %d, %d and %s",
cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout)
cfg.ClientRequestTimeout != 0 || cfg.ClientIdleTimeout != 0 {
t.Errorf("3G, off, off and off read as %d, %d, %s and %s",
cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout,
cfg.ClientIdleTimeout)
}
}
func TestRequestHeaderMaxBytesJustOver4K(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{clientHeaderMaxBytes: "4097"})
if cfg.ClientRequestHeaderMaxBytes != 4097 {
t.Errorf("4097 read as %d", cfg.ClientRequestHeaderMaxBytes)
}
}
func TestRequestHeaderMaxBytesRefusalNeverOffersOff(t *testing.T) {
t.Parallel()
for _, value := range []string{"32KB", "0", "4K", off} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(
environment{clientHeaderMaxBytes: value}.lookupEnv)
want := clientHeaderMaxBytes + `: "` + value +
`" is not a size of more than 4K, such as 32K`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
@@ -188,6 +262,15 @@ func TestRateLimitsOff(t *testing.T) {
}
}
func TestBanResponseCloseIsZero(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{banResponse: "close"})
if cfg.BanResponse != 0 {
t.Errorf("close read as %d, want 0", cfg.BanResponse)
}
}
func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
t.Parallel()
@@ -222,6 +305,8 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{denyNets, "198.51.100.0/24,"},
{clientRequestTimeout, "60"},
{clientRequestTimeout, ""},
{clientIdleTimeout, "0s"},
{clientIdleTimeout, "2 minutes"},
{clientResponseTimeout, "1y"},
{upstreamRequestTimeout, "-1s"},
{upstreamResponseTimeout, "0s"},
@@ -250,6 +335,15 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowedCountries, "uk"},
{allowedCountries, "zz"},
{allowedCountries, "de,germany"},
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
{maxBanDuration, off}, {maxBanDuration, "1w"},
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
{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.Parallel()
@@ -289,6 +383,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
upstreamURL: "http://127.0.0.1:8081",
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
clientRequestTimeout: "45s",
clientHeaderMaxBytes: "32K",
clientIdleTimeout: "120s",
clientResponseTimeout: "30m",
upstreamRequestTimeout: "60s",
upstreamResponseTimeout: "30m",
@@ -302,6 +398,15 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
rateLimitPerDay: "50000",
deniedCountries: "",
allowedCountries: "",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
maxBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -314,6 +419,8 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
if got.ListenAddr != want.ListenAddr ||
got.ClientRequestTimeout != want.ClientRequestTimeout ||
got.ClientRequestHeaderMaxBytes != want.ClientRequestHeaderMaxBytes ||
got.ClientIdleTimeout != want.ClientIdleTimeout ||
got.ClientResponseTimeout != want.ClientResponseTimeout ||
got.UpstreamRequestTimeout != want.UpstreamRequestTimeout ||
got.UpstreamResponseTimeout != want.UpstreamResponseTimeout ||
@@ -324,6 +431,28 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
got.RateLimitPerDay != want.RateLimitPerDay {
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
}
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans and the state files.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.BanResponse != want.BanResponse ||
got.LimitBanDuration != want.LimitBanDuration ||
got.LimitBanRepeatWindow != want.LimitBanRepeatWindow ||
got.MaxBanDuration != want.MaxBanDuration ||
got.MaxBans != want.MaxBans ||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
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.
+65 -13
View File
@@ -1,6 +1,7 @@
// 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
// 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
import (
@@ -12,6 +13,7 @@ import (
"log/slog"
"net/http"
"net/netip"
"slices"
"strings"
"sync"
"time"
@@ -76,7 +78,7 @@ type GeoJS struct {
httpClient *http.Client
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,
// and those it is being asked about.
waiting map[netip.Prefix]*wait
@@ -88,11 +90,14 @@ type GeoJS struct {
retryAt time.Time
}
// answer is what GeoJS said about a client: its country, "" when GeoJS
// cannot place it, and when GeoJS said so.
type answer struct {
country string
received time.Time
// Answer is what GeoJS said about a client, as lookups.json holds it: its
// country, "" when GeoJS cannot place it, when GeoJS said so, and when
// the answer was last used.
type Answer struct {
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.
@@ -107,7 +112,7 @@ type wait struct {
// New returns a GeoJS with no answer kept yet.
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 {
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
}
// 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
// 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
@@ -203,15 +249,19 @@ func (g *GeoJS) answerOrWait(
return "", w.asked
}
// kept returns client's answer, if one was received less than keepFor
// ago.
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
// and notes that it was used.
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
now := g.now()
kept, found := g.answers.Get(client)
if !found || g.now().Sub(kept.received) >= keepFor {
if !found || now.Sub(kept.Answered) >= keepFor {
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
@@ -293,7 +343,9 @@ func (g *GeoJS) keep(
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)
delete(g.waiting, client)
}
+96
View File
@@ -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)
}
}
+85
View File
@@ -0,0 +1,85 @@
package proxy
import (
"net/netip"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged
// with action.
func (rq *request) banResponse(action string) *refusal {
return &refusal{status: rq.h.config.BanResponse, action: action}
}
// banned reports whether a ban on a netblock the client is in refuses
// the request at now, and notes for the log line when that ban ends.
func (rq *request) banned(now time.Time) bool {
ban, banned := rq.h.ledger.Check(rq.client, now)
if banned {
rq.line.BanExpires = banExpires(ban)
}
return banned
}
// limitBroken counts the request for the rate limits at now, and reports
// whether it takes the client over one. Such a request bans the client's
// netblock, and sets the client's counters back to zero.
func (rq *request) limitBroken(now time.Time) bool {
group := clientGroup(rq.client)
hit, over := rq.h.limiter.Count(group, now)
if !over {
return false
}
netblock := rq.netblock()
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
},
// The histories count this request only once it has ended.
Requests: rq.h.limiter.Requests(netblock) + 1,
})
rq.h.limiter.Reset(group)
rq.line.LimitHit = hit.Window
rq.line.Offence = requestlog.OffenceLimit
rq.line.BanExpires = banExpires(ban)
return true
}
// netblock is the netblock a ban on the client covers: its IPv4 address,
// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup
// counts it in.
func (rq *request) netblock() netip.Prefix {
addr := rq.client.Unmap()
if addr.Is4() {
return netip.PrefixFrom(addr, rq.h.config.BanScopeV4Prefix).Masked()
}
return clientGroup(addr)
}
// banExpires is when ban ends, as the log line gives it: a time, or
// permanent.
func banExpires(ban bans.Ban) string {
if ban.Permanent() {
return "permanent"
}
return requestlog.FormatTime(ban.Expires)
}
+437
View File
@@ -0,0 +1,437 @@
package proxy_test
import (
"bufio"
"errors"
"io"
"maps"
"net/http"
"net/netip"
"slices"
"sync"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// otherClient is a client next to client.
otherClient = "203.0.113.10"
// userAgent is the user agent of every request a sender sends.
userAgent = "ban-test/1.0"
// permanent is the log line's ban_expires for a permanent ban.
permanent = "permanent"
)
func TestBrokenLimitBansTheClient(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{rateLimitPerMinute: "1"})
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
// The request over the limit of one a minute is refused, and bans the
// client for an hour, the default.
s.get(client, http.StatusOK, requestlog.ActionForward)
line := s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit ||
line.BanExpires != expires {
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
"want minute, limit and %s", line.LimitHit, line.Offence, line.BanExpires,
expires)
}
// Every request while the ban lasts is refused.
clk.advance(time.Hour - time.Second)
line = s.get(client, http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != expires || line.Offence != "" || line.LimitHit != "" {
t.Errorf("log line has ban_expires %q, offence %q and limit_hit %q, "+
"want %s and neither of the others", line.BanExpires, line.Offence,
line.LimitHit, expires)
}
// Once it ends, the client is let through.
clk.advance(time.Second)
s.get(client, http.StatusOK, requestlog.ActionForward)
}
func TestBanLengthsFollowTheSettings(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
limitBanDuration: "10m",
limitBanRepeatWindow: "1h",
maxBanDuration: "1h",
})
// breakLimit has client go over the limit of one a minute, and
// returns when the ban that makes ends.
breakLimit := func() string {
s.get(client, http.StatusOK, requestlog.ActionForward)
return s.get(client, http.StatusForbidden, requestlog.ActionRateLimited).BanExpires
}
wantExpires := func(got string, length time.Duration) {
t.Helper()
want := requestlog.FormatTime(clk.Now().Add(length))
if got != want {
t.Errorf("ban ends at %s, want %s", got, want)
}
}
// A first ban lasts SWWAF_LIMIT_BAN_DURATION; one within
// SWWAF_LIMIT_BAN_REPEAT_WINDOW after it ended, three times as long.
wantExpires(breakLimit(), 10*time.Minute)
clk.advance(10*time.Minute + time.Hour)
wantExpires(breakLimit(), 30*time.Minute)
// Later than that, SWWAF_LIMIT_BAN_DURATION again.
clk.advance(30*time.Minute + time.Hour + time.Second)
wantExpires(breakLimit(), 10*time.Minute)
clk.advance(10 * time.Minute)
wantExpires(breakLimit(), 30*time.Minute)
// 90 minutes would be longer than SWWAF_MAX_BAN_DURATION: the ban is
// permanent.
clk.advance(30 * time.Minute)
got := breakLimit()
if got != permanent {
t.Errorf("ban ends at %s, want a permanent one", got)
}
clk.advance(365 * 24 * time.Hour)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{rateLimitPerDay: "2"})
// The third request in a day is over the limit of two, and bans the
// client for an hour.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
for range 3 {
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
// Later the same day the client has its whole allowance again: the
// ban set its counters back to zero, and the requests it refused were
// not counted for the rate limits, only in its notes.
clk.advance(time.Hour)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(banned) != 2 || banned[0].Notes.Refused != 3 {
t.Errorf("bans %+v, want two, the first with 3 requests refused", banned)
}
}
func TestBanCoversTheClientsNetblock(t *testing.T) {
t.Parallel()
// In the IPv4 cases, client breaks the limit; these two are next to it.
const (
allowed = "203.0.113.60" // in SWWAF_ALLOW_NETS
exempt = "203.0.113.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
)
for _, tc := range []struct {
name string
env map[string]string
breaker string // the client that breaks the limit
refused []string
let []string // let through
}{
{
"an IPv4 address, by default", nil, client,
nil, []string{otherClient, exempt},
},
{
"the IPv4 netblock SWWAF_BAN_SCOPE_V4_PREFIX sets",
map[string]string{banScopeV4Prefix: "24"}, client,
[]string{otherClient, exempt}, []string{"203.0.112.9", allowed},
},
{
"an IPv6 /64", nil, "2001:db8:5::1",
[]string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{
rateLimitPerMinute: "1",
allowNets: allowed,
rateLimitExemptNets: exempt,
}
maps.Copy(env, tc.env)
s, _, _ := startWithClock(t, "", env)
s.get(tc.breaker, http.StatusOK, requestlog.ActionForward)
s.get(tc.breaker, http.StatusForbidden, requestlog.ActionRateLimited)
for _, sent := range tc.refused {
s.get(sent, http.StatusForbidden, requestlog.ActionBanned)
}
for _, sent := range tc.let {
s.get(sent, http.StatusOK, requestlog.ActionForward)
}
})
}
}
func TestBannedClientIsRefusedBeforeItsCountryIsLookedUp(t *testing.T) {
t.Parallel()
geojsURL, asked := startGeoJS(t)
s, _, _ := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1",
banScopeV4Prefix: "24",
deniedCountries: "kp",
})
// fromDE's ban covers otherClient, which is refused unasked about.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
line := s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned)
if line.Country != "" {
t.Errorf("log line has country %q, want none", line.Country)
}
if !slices.Equal(asked(), []string{fromDE}) {
t.Errorf("GeoJS was asked about %v, want %s alone", asked(), fromDE)
}
}
func TestBanResponseAnswersEveryRefusalButTheSizeLimits(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
for _, tc := range []struct {
setting string // "" leaves SWWAF_BAN_RESPONSE at its default
status int // 0 is the connection closed without an answer
}{
{"", http.StatusForbidden},
{"403", http.StatusForbidden},
{"429", http.StatusTooManyRequests},
{"close", 0},
} {
t.Run(banResponse+"="+tc.setting, func(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
env := map[string]string{
rateLimitPerMinute: "1",
denyNets: denied,
deniedCountries: "kp",
}
if tc.setting != "" {
env[banResponse] = tc.setting
}
s, _, _ := startWithClock(t, geojsURL, env)
s.get(denied, tc.status, requestlog.ActionDenied)
s.get(fromKP, tc.status, requestlog.ActionCountryDenied)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, tc.status, requestlog.ActionRateLimited)
s.get(fromDE, tc.status, requestlog.ActionBanned)
})
}
}
func TestBanNotes(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1",
deniedCountries: "kp",
})
start := clk.Now()
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.request(fromDE, "/repo/commits?page=2",
http.StatusForbidden, requestlog.ActionRateLimited)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
netblock := netip.MustParsePrefix(fromDE + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: start.Add(time.Hour),
Notes: bans.Notes{
Country: "DE",
Limit: 1,
Window: minute,
Count: 2,
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: "/repo/commits?page=2",
Status: http.StatusForbidden,
UserAgent: userAgent,
},
// The one let through, the one that broke the limit and the two
// refused under the ban.
Requests: 4,
Refused: 2,
EarlierBans: 0,
},
}
ledger := server.Ledger
got := ledger.Bans(netblock)
if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
// The next ban counts this one among the earlier.
clk.advance(time.Hour)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
got = ledger.Bans(netblock)
if len(got) != 2 || got[1].Notes.EarlierBans != 1 {
t.Errorf("bans %+v, want two, the second with one earlier ban", got)
}
}
func TestMaxBansDropsTheBanOfTheNetblockSeenLongestAgo(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
maxBans: "1",
})
// One ban is held, so otherClient's ban drops client's.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned)
}
// clock is the time a test sets, by which smallwebwaf counts requests and
// makes bans.
type clock struct {
mu sync.Mutex
now time.Time
}
// Now tells the time.
func (c *clock) Now() time.Time {
c.mu.Lock()
defer c.mu.Unlock()
return c.now
}
// advance moves the clock on by d.
func (c *clock) advance(d time.Duration) {
c.mu.Lock()
defer c.mu.Unlock()
c.now = c.now.Add(d)
}
// startWithClock starts smallwebwaf in front of an app that answers 200,
// with the settings in env on top of trusting localhost's
// X-Forwarded-For, clients' countries looked up at geojsURL, and a clock
// set to midnight, the start of a bucket in every window.
func startWithClock(
t *testing.T, geojsURL string, env map[string]string,
) (*sender, *clock, *proxy.Server) {
t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{trustedProxies: trustLocalhost}
maps.Copy(settings, env)
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, clk.Now, settings)
return &sender{t: t, addr: addr, out: out}, clk, server
}
// sender sends requests to smallwebwaf one after another, each on a
// connection of its own, and checks each one's answer and log line. They
// must be the only requests smallwebwaf is sent, since the log lines are
// matched to them in order.
type sender struct {
t *testing.T
addr string
out *output
sent int
}
// get sends a GET request for / from the client at from.
func (s *sender) get(from string, status int, action string) logLine {
s.t.Helper()
return s.request(from, "/", status, action)
}
// request sends a GET request for path from the client at from, as
// X-Forwarded-For names it, and checks that its answer and its log line
// have status, 0 for the connection closed without an answer, and that
// the line has action. It returns the log line.
func (s *sender) request(from, path string, status int, action string) logLine {
s.t.Helper()
conn := dial(s.t, s.addr)
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n\r\n")
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
s.t.Fatalf("set read deadline: %v", err)
}
got := 0
res, err := http.ReadResponse(bufio.NewReader(conn), nil)
switch {
case err == nil:
got = readAnswer(res).status
case !errors.Is(err, io.ErrUnexpectedEOF):
s.t.Fatalf("read response: %v", err)
}
_ = conn.Close()
if got != status {
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, got,
status)
}
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
s.sent++
wantLine(s.t, line, status, action)
return line
}
+103
View File
@@ -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{}
}
+37 -26
View File
@@ -297,7 +297,7 @@ func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) {
}
}
func TestServerHasTheFixedLimits(t *testing.T) {
func TestServerHasTheDefaultLimits(t *testing.T) {
t.Parallel()
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
@@ -319,37 +319,48 @@ func TestServerHasTheFixedLimits(t *testing.T) {
}
}
func TestRefusesHeadersOver32KiB(t *testing.T) {
func TestRefusesHeadersOverTheLimit(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, _ := startProxy(t, app.URL, nil)
// size counts every byte of the request: the request line, the
// headers and the blank line that ends them.
const (
start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: "
end = "\r\n\r\n"
)
for _, tc := range []struct {
size int
want int
name string
env map[string]string
limit int
}{
{size: 32 << 10, want: http.StatusOK},
{size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge},
{"by default", nil, 32 << 10},
{"as set", map[string]string{clientHeaderMaxBytes: "8K"}, 8 << 10},
} {
conn := dial(t, addr)
send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end)
wantStatus(t, readResponse(t, conn), tc.want)
}
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
if calls.Load() != 1 {
t.Errorf("the app was called %d times, want once", calls.Load())
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, _ := startProxy(t, app.URL, tc.env)
// size counts every byte of the request: the request line,
// the headers and the blank line that ends them.
const (
start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: "
end = "\r\n\r\n"
)
for _, sent := range []struct{ size, want int }{
{tc.limit, http.StatusOK},
{tc.limit + 1, http.StatusRequestHeaderFieldsTooLarge},
} {
conn := dial(t, addr)
send(t, conn,
start+strings.Repeat("a", sent.size-len(start)-len(end))+end)
wantStatus(t, readResponse(t, conn), sent.want)
}
if calls.Load() != 1 {
t.Errorf("the app was called %d times, want once", calls.Load())
}
})
}
}
+63 -38
View File
@@ -10,25 +10,13 @@ import (
"net/http"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The request line and headers a client may send, and how long a
// kept-open client connection may wait for its next request, are fixed
// rather than settings. The limit on the request line and headers is
// 32 KiB, but Go's server reads 4 KiB past its MaxHeaderBytes before it
// refuses, so MaxHeaderBytes is set 4 KiB lower. The idle time is longer
// than the 90 seconds after which traefik closes a connection it is not
// using, so traefik never sends a request on a connection smallwebwaf is
// closing.
const (
requestHeaderMaxBytes = 32<<10 - 4<<10
clientIdleTimeout = 120 * time.Second
)
// How smallwebwaf keeps connections to the app open between requests.
const (
appIdleConns = 100
@@ -49,39 +37,71 @@ type Params struct {
// GeoJSURL is where clients' countries are looked up, normally
// lookup.URL. GeoJS is asked only while a country list is set.
GeoJSURL string
// Now tells the time by which requests are counted for the rate
// 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
}
// 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
// through the proxy. Go's server itself refuses headers over 32 KiB, with
// 431, closes a connection idle for 120 seconds, and applies
// through the proxy. Go's server itself refuses a request line and
// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
// 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)
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{
Addr: params.Config.ListenAddr,
Handler: &handler{
config: params.Config,
requestLog: params.RequestLog,
processLog: params.ProcessLog,
errorLog: errorLog,
transport: newTransport(),
limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay,
}),
geojs: lookup.New(lookup.Params{
URL: params.GeoJSURL,
Now: time.Now,
ProcessLog: params.ProcessLog,
}),
return &Server{
Server: &http.Server{
Addr: params.Config.ListenAddr,
Handler: h,
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
// Off is an IdleTimeout of 0, which Go's server replaces with
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
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,
},
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
IdleTimeout: clientIdleTimeout,
MaxHeaderBytes: requestHeaderMaxBytes,
ErrorLog: errorLog,
Ledger: h.ledger,
Limiter: h.limiter,
GeoJS: h.geojs,
}
}
@@ -93,7 +113,9 @@ type handler struct {
processLog *slog.Logger
errorLog *log.Logger
transport http.RoundTripper
now func() time.Time
limiter *ratelimit.Limiter
ledger *bans.Ledger
geojs *lookup.GeoJS
}
@@ -125,6 +147,9 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
// Once the request has ended, before its log line is written.
defer rq.addToHistory()
refused := rq.check(r.Context())
if refused != nil {
rq.answer(*refused)
+24 -1
View File
@@ -45,6 +45,8 @@ var shortTimeoutSetting = shortTimeout.String()
// The settings the tests set.
const (
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
@@ -55,8 +57,15 @@ const (
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
)
// output collects what smallwebwaf writes on stdout.
@@ -181,6 +190,19 @@ func startProxyWithGeoJS(
) (string, *output) {
t.Helper()
addr, out, _ := startProxyWithClock(t, appURL, geojsURL, time.Now, env)
return addr, out
}
// startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well.
func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *proxy.Server) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
maps.Copy(settings, env)
@@ -199,6 +221,7 @@ func startProxyWithGeoJS(
RequestLog: out,
ProcessLog: requestlog.NewProcessLogger(out),
GeoJSURL: geojsURL,
Now: now,
})
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
@@ -214,7 +237,7 @@ func startProxyWithGeoJS(
_ = server.Close()
})
return listener.Addr().String(), out
return listener.Addr().String(), out, server
}
// newClient returns an HTTP client that sends requests as they are made,
+12 -8
View File
@@ -8,7 +8,11 @@ import (
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
// minute is the window of SWWAF_RATE_LIMIT_PER_MINUTE, as a log line's
// limit_hit names it.
const minute = "minute"
func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
t.Parallel()
var calls atomic.Int32
@@ -24,19 +28,19 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
const otherClient = "203.0.113.10"
// With a limit of one request a minute, a client's second request is
// refused. A client is one IPv4 address, or one IPv6 /64; an IPv4
// address in IPv6 form is that IPv4 address.
// refused, with 403 by default. A client is one IPv4 address, or one
// IPv6 /64; an IPv4 address in IPv6 form is that IPv4 address.
requests := []struct {
client string // as X-Forwarded-For names it
logged string // as the log line's client_ip names it
want int
}{
{client, client, http.StatusOK},
{client, client, http.StatusTooManyRequests},
{client, client, http.StatusForbidden},
{otherClient, otherClient, http.StatusOK},
{"::ffff:" + otherClient, otherClient, http.StatusTooManyRequests},
{"::ffff:" + otherClient, otherClient, http.StatusForbidden},
{"2001:db8::1", "2001:db8::1", http.StatusOK},
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusTooManyRequests},
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusForbidden},
{"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK},
}
@@ -53,9 +57,9 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
if sent.want == http.StatusOK {
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
} else {
wantLine(t, line, http.StatusTooManyRequests, requestlog.ActionRateLimited)
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
if line.LimitHit != "minute" {
if line.LimitHit != minute {
t.Errorf("log line has limit_hit %q, want minute", line.LimitHit)
}
}
+44 -23
View File
@@ -12,6 +12,7 @@ import (
"sync/atomic"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
@@ -21,7 +22,8 @@ const flushAfterEachWrite time.Duration = -1
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
// the status the client is answered if the response has not started yet,
// and the action the log line names.
// 0 to close the connection without an answer, and the action the log
// line names.
type refusal struct {
status int
action string
@@ -105,39 +107,33 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
// is known, before its body is read or anything reaches the app. It
// returns nil to let the request through. A client in SWWAF_ALLOW_NETS
// skips every check but the size limit. For any other client,
// SWWAF_DENY_NETS comes first, so that a client it refuses is not looked
// up, and then the country lists; a request either refuses is not counted
// for the rate limits. Then come the rate limits, unless the client is in
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
// client either refuses is not looked up, and then the country lists; a
// request any of them refuses is not counted for the rate limits. Then
// come the rate limits, unless the client is in
// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted,
// one refused for its size too. ctx is the request's own context.
// one refused for its size too. Every refusal but the size limit's is
// answered with SWWAF_BAN_RESPONSE. ctx is the request's own context.
func (rq *request) check(ctx context.Context) *refusal {
cfg := rq.h.config
allowed := isInside(rq.client, cfg.AllowNets)
exempt := isInside(rq.client, cfg.RateLimitExemptNets)
now := rq.h.now()
if !allowed && isInside(rq.client, cfg.DenyNets) {
return &refusal{
status: http.StatusForbidden,
action: requestlog.ActionDenied,
}
return rq.banResponse(requestlog.ActionDenied)
}
if !allowed && rq.banned(now) {
return rq.banResponse(requestlog.ActionBanned)
}
if !allowed && rq.countryDenied(ctx) {
return &refusal{
status: http.StatusForbidden,
action: requestlog.ActionCountryDenied,
}
return rq.banResponse(requestlog.ActionCountryDenied)
}
if !allowed && !isInside(rq.client, cfg.RateLimitExemptNets) {
limitHit := rq.h.limiter.Count(clientGroup(rq.client), rq.start)
if limitHit != "" {
rq.line.LimitHit = limitHit
return &refusal{
status: http.StatusTooManyRequests,
action: requestlog.ActionRateLimited,
}
}
if !allowed && !exempt && rq.limitBroken(now) {
return rq.banResponse(requestlog.ActionRateLimited)
}
maxBytes := cfg.RequestMaxBytes
@@ -250,6 +246,13 @@ func (rq *request) answer(r refusal) {
return // too late to answer: the connection can only be cut
}
if r.status == 0 {
// SWWAF_BAN_RESPONSE is close. This panic has Go's server close
// the connection without an answer, and log nothing; the log line
// is still written as the handler returns.
panic(http.ErrAbortHandler)
}
// A client found too slow is read no more; any other may go on
// sending until its time is up, so that Go's server can read the
// rest of the body and end the request cleanly.
@@ -316,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
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
func (rq *request) clientRequestDeadline() time.Time {
+3 -3
View File
@@ -78,7 +78,7 @@ func TestRequestFromAllowNetsIsNotCounted(t *testing.T) {
{listedAddr, http.StatusOK, requestlog.ActionForward},
{listedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
})
}
@@ -133,7 +133,7 @@ func TestRequestRefusedByDenyNetsIsNotCounted(t *testing.T) {
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
})
}
@@ -156,7 +156,7 @@ func TestRateLimitExemptNetsAreNeitherCountedNorRefused(t *testing.T) {
{listedAddr, http.StatusOK, requestlog.ActionForward},
{listedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
{fromKP, http.StatusForbidden, requestlog.ActionCountryDenied},
})
}
+23
View File
@@ -273,3 +273,26 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
wantTimedOut(t, start)
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
}
func TestClosesAnIdleConnection(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, _ := startProxy(t, app.URL, map[string]string{
clientIdleTimeout: shortTimeoutSetting,
})
// The idle time starts once the answer is sent, so after start.
start := time.Now()
conn := dial(t, addr)
send(t, conn, "GET / HTTP/1.1\r\nHost: app\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusOK)
// The read deadline readResponse set still bounds this read.
_, err := conn.Read(make([]byte, 1))
if !errors.Is(err, io.EOF) {
t.Fatalf("read on the idle connection: %v, want it closed", err)
}
wantTimedOut(t, start)
}
+116
View File
@@ -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{}
}
+273 -47
View File
@@ -1,11 +1,15 @@
// Package ratelimit counts each client's requests over a minute, an hour
// and a day, as the "Counting method" section of SPEC.md describes, and
// tells when a request takes a client over a rate limit. The counts are
// kept in memory only, for at most 20,000 clients.
// Package ratelimit keeps the table of clients: each client's requests
// counted over a minute, an hour and a day, as the "Counting method"
// section of SPEC.md describes, which tell when a request takes the client
// 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
import (
"net/http"
"net/netip"
"slices"
"sync"
"time"
@@ -13,7 +17,8 @@ import (
)
// 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 day = 24 * time.Hour
@@ -26,20 +31,94 @@ type Limits struct {
PerDay int64
}
// Limiter counts each client's requests against the limits. It is safe
// for concurrent use.
// Limiter counts each client's requests against the limits, and keeps
// its history. It is safe for concurrent use.
type Limiter struct {
// windows are the minute, the hour and the day, in the order of
// Client.buckets.
windows [3]window
mu sync.Mutex
// clients holds each client's buckets, one pair for each of windows,
// in the same order.
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
mu sync.Mutex
clients *simplelru.LRU[netip.Prefix, *Client]
}
// 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.
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 {
panic(err) // NewLRU fails only for a size below one
}
@@ -54,30 +133,168 @@ func New(limits Limits) *Limiter {
}
}
// Hit is a request that takes a client over a rate limit.
type Hit struct {
// Window is "minute", "hour" or "day".
Window string
// Limit is the window's limit.
Limit int64
// Requests is the client's requests counted in the window, this one
// included.
Requests float64
}
// Count counts a request from client at now, in every window, whether or
// not it is refused. It returns the window whose limit the request takes
// the client over, "minute", "hour" or "day", the shortest if it is over
// several, or "" if it is within every limit.
func (l *Limiter) Count(client netip.Prefix, now time.Time) string {
// not it is refused. It reports whether the request takes the client over
// a limit, and the window whose limit it goes over, the shortest if it is
// over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
counts, seen := l.clients.Get(client)
if !seen {
counts = &[3]buckets{}
l.clients.Add(client, counts)
}
var hit Hit
limitHit := ""
for i, b := range l.get(client).buckets() {
w := l.windows[i]
for i, w := range l.windows {
requests := counts[i].add(now, w.length)
if limitHit == "" && w.limit > 0 && requests > float64(w.limit) {
limitHit = w.name
requests := b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
}
}
return limitHit
return hit, hit.Window != ""
}
// Reset sets client's counts in every window back to zero. Its history
// keeps its totals.
func (l *Limiter) Reset(client netip.Prefix) {
l.mu.Lock()
defer l.mu.Unlock()
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
@@ -88,14 +305,6 @@ type window struct {
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
// 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
@@ -106,27 +315,44 @@ type buckets struct {
// 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
// would keep its full weight until the clock caught up.
func (b *buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.start.Add(-time.Second)) {
*b = buckets{}
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.Start.Add(-time.Second)) {
*b = Buckets{}
}
start := now.Truncate(length)
if start.After(b.start) {
if start.Equal(b.start.Add(length)) {
b.previous = b.current
if start.After(b.Start) {
if start.Equal(b.Start.Add(length)) {
b.Previous = b.Current
} else {
b.previous = 0
b.Previous = 0
}
b.start = start
b.current = 0
b.Start = start
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)
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++
}
}
+49 -3
View File
@@ -54,6 +54,52 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
}
}
func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
_, over := limiter.Count(client, start)
if over {
t.Fatal("a request within the limit is over it")
}
}
// Over both limits; the minute's is named, with the four requests.
hit, over := limiter.Count(client, start)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want {
t.Errorf("request over the limit gives %+v and %t, want %+v and true",
hit, over, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
wantCount(t, limiter, client, start, "")
}
wantCount(t, limiter, client, start, minute)
limiter.Reset(client)
// At the same moment, the client has its whole allowance again.
for range limit {
wantCount(t, limiter, client, start, "")
}
wantCount(t, limiter, client, start, minute)
}
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
t.Parallel()
@@ -192,9 +238,9 @@ func wantCount(
) {
t.Helper()
got := limiter.Count(client, now)
if got != want {
hit, _ := limiter.Count(client, now)
if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), got, want)
client, now.Format(time.RFC3339), hit.Window, want)
}
}
+122
View File
@@ -0,0 +1,122 @@
package ratelimit_test
import (
"net/netip"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
func TestSnapshotListsTheClientsByAddress(t *testing.T) {
t.Parallel()
want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"}
limiter := ratelimit.New(ratelimit.Limits{})
for _, i := range []int{2, 3, 0, 1} {
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
}
snapshot := limiter.Snapshot()
got := make([]string, 0, len(snapshot))
for _, c := range snapshot {
got = append(got, c.Client.String())
}
if !slices.Equal(got, want) {
t.Errorf("snapshot %v, want %v", got, want)
}
counted := ratelimit.Buckets{Start: midnight(), Current: 1}
if snapshot[0].Minute != counted || snapshot[0].Day != counted {
t.Errorf("buckets %+v and %+v, want %+v", snapshot[0].Minute, snapshot[0].Day,
counted)
}
}
func TestLoadedCountsCarryOn(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
before := ratelimit.New(ratelimit.Limits{PerHour: limit})
for range limit {
wantCount(t, before, client, start, "")
}
// Loaded into a new limiter, as across a restart, the client has no
// fresh allowance.
later := start.Add(time.Minute)
after := ratelimit.New(ratelimit.Limits{PerHour: limit})
after.Load(before.Snapshot(), later)
wantCount(t, after, client, later, hour)
}
func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client {
t.Helper()
after := ratelimit.New(ratelimit.Limits{})
after.Load(limiter.Snapshot(), now)
return after.Snapshot()[0]
}
// Two minutes on, the window that ends then covers neither of the
// minute's buckets, which are emptied; the hour's and the day's stay,
// and so does the history.
got := loaded(start.Add(2 * time.Minute))
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
got.Day.Current != 1 || got.History.Requests != 1 {
t.Errorf("loaded two minutes on as %+v", got)
}
// A moment before, the window still covers some of the earlier one.
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
if got.Minute.Current != 1 {
t.Errorf("loaded just under two minutes on with minute buckets %+v",
got.Minute)
}
}
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
t.Parallel()
const maxClients = 20000
// clients.json lists the clients by address. Here each was last seen
// a second before the one listed before it, so the last listed is the
// one seen longest ago, and the one dropped.
clients := make([]ratelimit.Client, maxClients+1)
addr := netip.MustParseAddr("10.0.0.0")
for i := range clients {
clients[i].Client = netip.PrefixFrom(addr, addr.BitLen())
clients[i].History.LastSeen = midnight().Add(-time.Duration(i) * time.Second)
addr = addr.Next()
}
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Load(clients, midnight())
got := limiter.Snapshot()
if len(got) != maxClients || got[0].Client != clients[0].Client ||
got[maxClients-1].Client != clients[maxClients-1].Client {
t.Errorf("%d clients kept, from %s to %s; want %d, from %s to %s",
len(got), got[0].Client, got[len(got)-1].Client, maxClients,
clients[0].Client, clients[maxClients-1].Client)
}
}
+12 -1
View File
@@ -24,8 +24,10 @@ const (
// for, or whose answer could not be passed on.
ActionUpstreamError = "upstream_error"
// ActionRateLimited is a request refused because it took its client
// over a rate limit, or came while the client was over one.
// over a rate limit, which bans the client.
ActionRateLimited = "rate_limited"
// ActionBanned is a request refused because a ban covers its client.
ActionBanned = "banned"
// ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS.
ActionDenied = "denied"
@@ -36,6 +38,10 @@ const (
ActionAdmin = "admin"
)
// OffenceLimit is the offence a request line names for a request that
// broke a rate limit.
const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
@@ -64,6 +70,11 @@ type Line struct {
// LimitHit is the window whose rate limit the request went over:
// minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"`
// Offence is the offence the request was held as, OffenceLimit.
Offence string `json:"offence,omitempty"`
// BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent".
BanExpires string `json:"ban_expires,omitempty"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"`
// DurationTotal and DurationUpstreamTotal are in milliseconds.
+2 -1
View File
@@ -50,7 +50,8 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
}
unset := []string{
"upstream_status", "limit_hit", "aborted", "duration_upstream_total",
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted",
"duration_upstream_total",
}
for _, name := range unset {
_, present := fields[name]
+6 -4
View File
@@ -24,12 +24,14 @@ func TestHealthCheck(t *testing.T) {
out := &output{}
exited := make(chan int, 1)
settings := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
stateDir: t.TempDir(),
}
go func() {
exited <- run(ctx, map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
}, out)
exited <- run(ctx, settings, out)
}()
addr, _ := out.line(t, "msg", "starting")["address"].(string)
+64 -16
View File
@@ -1,6 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// serves requests until it is told to stop, and then stops in an orderly
// way.
// Package smallwebwaf runs the smallwebwaf process: it reads the settings
// and the state files, serves requests until it is told to stop, and then
// stops in an orderly way, writing the state files.
package smallwebwaf
import (
@@ -19,6 +19,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
// 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
// returns the process's exit status, 1 when smallwebwaf cannot start.
// Run reads the settings and the state files, then serves requests until
// ctx is done. It returns the process's exit status, 1 when smallwebwaf
// cannot start.
func Run(ctx context.Context, params Params) int {
processLog := requestlog.NewProcessLogger(params.Stdout)
@@ -67,6 +69,33 @@ func Run(ctx context.Context, params Params) int {
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)
if err != nil {
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
@@ -75,26 +104,20 @@ func Run(ctx context.Context, params Params) int {
return 1
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
})
processLog.Info("starting",
"version", params.Version,
"address", listener.Addr().String(),
"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
// requests in progress shutdownTimeout to finish.
// serve serves requests on listener, and writes the state files as they
// are due, until ctx is done. Then it gives the requests in progress
// shutdownTimeout to finish, and writes every state file.
func serve(
ctx context.Context, server *http.Server, listener net.Listener,
processLog *slog.Logger,
files *state.Files, processLog *slog.Logger,
) int {
served := make(chan error, 1)
@@ -102,6 +125,16 @@ func serve(
served <- server.Serve(listener)
}()
writing, stopWriting := context.WithCancel(ctx)
defer stopWriting()
written := make(chan struct{})
go func() {
files.Run(writing)
close(written)
}()
select {
case err := <-served:
processLog.Error("serving failed", "error", err.Error())
@@ -131,6 +164,21 @@ func serve(
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")
return 0
+277 -31
View File
@@ -8,6 +8,8 @@ import (
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing"
@@ -24,9 +26,13 @@ const (
// testVersion is the version the tests give smallwebwaf.
testVersion = "test"
// localhost is where the tests listen.
localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL"
localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR"
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.
@@ -69,11 +75,19 @@ func (o *output) line(t *testing.T, key, value string) map[string]any {
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
}
// 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
// returns its exit status.
func run(ctx context.Context, env map[string]string, out *output) int {
@@ -121,7 +135,10 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
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 {
t.Errorf("exit status %d, want 1", status)
}
@@ -132,11 +149,8 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
func TestServesUntilToldToStop(t *testing.T) {
t.Parallel()
app := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "hello from the app")
}))
defer app.Close()
appURL := startApp(t)
dir := t.TempDir()
ctx, stop := context.WithCancel(t.Context())
out := &output{}
@@ -145,12 +159,13 @@ func TestServesUntilToldToStop(t *testing.T) {
go func() {
exited <- run(ctx, map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
upstreamURL: appURL,
stateDir: dir,
}, out)
}()
starting := out.line(t, "msg", "starting")
wantStartingLine(t, starting, app.URL)
wantStartingLine(t, starting, appURL, dir)
addr, _ := starting["address"].(string)
wantGreeting(t, "http://"+addr+"/")
@@ -170,30 +185,207 @@ func TestServesUntilToldToStop(t *testing.T) {
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
// 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()
settings, _ := line["settings"].(map[string]any)
want := map[string]any{
listenAddr: localhost + ":0",
upstreamURL: appURL,
"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_RESPONSE_TIMEOUT": "30m",
"SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s",
"SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m",
"SWWAF_REQUEST_MAX_BYTES": "100M",
"SWWAF_RESPONSE_MAX_BYTES": "5G",
"SWWAF_ALLOW_NETS": "",
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
"SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
"SWWAF_RATE_LIMIT_PER_DAY": "50000",
"SWWAF_DENIED_COUNTRIES": "",
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
listenAddr: localhost + ":0",
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_CLIENT_REQUEST_TIMEOUT": "60s",
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
"SWWAF_CLIENT_IDLE_TIMEOUT": "120s",
"SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m",
"SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s",
"SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m",
"SWWAF_REQUEST_MAX_BYTES": "100M",
"SWWAF_RESPONSE_MAX_BYTES": "5G",
"SWWAF_ALLOW_NETS": "",
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
"SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
rateLimitPerDay: "50000",
"SWWAF_DENIED_COUNTRIES": "",
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
"SWWAF_BAN_RESPONSE": "403",
"SWWAF_LIMIT_BAN_DURATION": "1h",
"SWWAF_LIMIT_BAN_REPEAT_WINDOW": "24h",
"SWWAF_MAX_BAN_DURATION": "7d",
"SWWAF_MAX_BANS": "5000",
"SWWAF_BAN_SCOPE_V4_PREFIX": "32",
}
for name, value := range want {
@@ -228,7 +420,61 @@ func wantGreeting(t *testing.T, url string) {
body, err := io.ReadAll(res.Body)
_ = 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)
}
}
// 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)
}
}
+492
View File
@@ -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())
}
+635
View File
@@ -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
View File
@@ -1,12 +1,14 @@
#!/bin/sh
# 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
# health check passes, that a request is served through smallwebwaf,
# that `sv stop` stops smallwebwaf in order, and that `docker stop`
# stops the container without having to kill it. The container 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.
# deploy/example-app, then run the app's container with a volume for the
# state files and check that the health check passes, that a request is
# served through smallwebwaf, that a second one in a minute bans the
# client, that `sv stop` stops smallwebwaf in order, that `docker stop`
# stops the container without having to kill it, and that a new
# container on the same volume still refuses the banned client. The
# containers, the volume and both images are removed however the script
# ends. Building the app needs network access, for nixpkgs' binary cache.
# script/check does not run this.
set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -18,9 +20,11 @@ NAME="$("$SCRIPT_DIR/projectname")-example-$$"
IMAGE="$NAME-base"
APP_IMAGE="$NAME-app"
CONTAINER="$NAME"
VOLUME="$NAME-state"
cleanup() {
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
}
@@ -53,6 +57,26 @@ logged() {
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() {
cd "$ROOT"
trap cleanup EXIT
@@ -62,18 +86,20 @@ main() {
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
-t "$APP_IMAGE" deploy/example-app
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
"$APP_IMAGE" >/dev/null
wait_for "the health check did not pass" healthy
docker volume create "$VOLUME" >/dev/null
start_container
echo "example-app: the health check passes"
address="$(docker port "$CONTAINER" 8080/tcp)"
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
fail "no answer on port 8080"
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
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"
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 ||
fail "sv stop smallwebwaf failed"
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
@@ -83,6 +109,13 @@ main() {
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
[ "$status" = 0 ] || fail "docker stop left exit status $status"
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 "$@"
+7 -1
View File
@@ -1,6 +1,7 @@
#!/bin/sh
# 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
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -8,6 +9,11 @@ ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)"
main() {
"$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"
}
+6 -2
View File
@@ -2,10 +2,14 @@
set -euo pipefail
# runit's run script for smallwebwaf, run again whenever smallwebwaf
# exits; the wait spaces out the restarts. exec, so that the signal
# `sv stop` sends reaches smallwebwaf itself.
# exits; the wait spaces out the restarts. The state directory and every
# 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() {
sleep 1
chown -R smallwebwaf:smallwebwaf "${SWWAF_STATE_DIR:-/var/lib/smallwebwaf}"
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
}