Compare commits
14
Commits
fe69d20a12
..
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bff65f4e2f | ||
|
|
ee9ba08a8a | ||
|
|
0797e5def2 | ||
|
|
e77dfb6891 | ||
|
|
74bdc6a449 | ||
|
|
808e69f442 | ||
|
|
6ec52e5b87 | ||
|
|
cff385af41 | ||
|
|
234c5eac60 | ||
|
|
68f687cb0c | ||
|
|
df2c5042d2 | ||
|
|
73ca94f850 | ||
|
|
0f85c9ae07 | ||
|
|
50df9ee36e |
@@ -162,6 +162,15 @@ RUN groupadd --system --gid 65532 smallwebwaf \
|
|||||||
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
|
--gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \
|
||||||
smallwebwaf
|
smallwebwaf
|
||||||
|
|
||||||
|
# The state files' directory, SWWAF_STATE_DIR by default, where a volume
|
||||||
|
# is mounted to keep them across deploys. The run script gives it to the
|
||||||
|
# smallwebwaf user at each start.
|
||||||
|
RUN mkdir /var/lib/smallwebwaf
|
||||||
|
|
||||||
|
# The default rule file, in SWWAF_RULES_DIR by default, where an app's
|
||||||
|
# Dockerfile can copy rule files of its own beside it.
|
||||||
|
COPY share/rules.d/00-default.rules /etc/smallwebwaf/rules.d/00-default.rules
|
||||||
|
|
||||||
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
|
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
|
||||||
# looks too.
|
# looks too.
|
||||||
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
|
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
|
||||||
|
|||||||
@@ -13,15 +13,30 @@ JSON log line for every request.
|
|||||||
|
|
||||||
Status: the first two milestones are built
|
Status: the first two milestones are built
|
||||||
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
|
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
|
||||||
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are the static lists,
|
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are nine parts of
|
||||||
which come next in the build order. `smallwebwaf` passes each request to the app
|
milestone 3: the static lists, the bans that broken rate limits lead to, the ban
|
||||||
|
ledger with the bans you make, keep and lift, the JSON state files with your
|
||||||
|
edits taken in while it runs and the paths the rate limits do not count, which
|
||||||
|
come next in the build order, `observe` mode and the rest of the request log's
|
||||||
|
fields, which come a little later, and the metrics endpoint and the header size
|
||||||
|
and the idle time as settings, which come last in it. So are two parts of the
|
||||||
|
stage after it: the rule files, the first part, with the bans for a clear sign
|
||||||
|
of attack, and remote log sending. `smallwebwaf` passes each request to the app
|
||||||
and the app's answer back, unchanged, within its timeouts and size limits, works
|
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
|
out each client's address, bans a client that sends too many requests, not
|
||||||
from a country you refuse or from a network you refuse, lets the networks you
|
counting those for the paths you choose, refuses a client that comes from a
|
||||||
choose through, and writes a JSON log line for every request. It comes as the
|
country you refuse or from a network you refuse, lets the networks you choose
|
||||||
image the app's own image is built on. The rest of the design comes after that,
|
through, checks each request against the rule files and bans a client whose
|
||||||
in the order of the build order in [`SPEC.md`](SPEC.md). The survey of existing
|
request is a clear sign of attack, keeps its bans, each client's counters and
|
||||||
tools that led to the design is in [`EVALUATION.md`](EVALUATION.md).
|
history, and GeoJS's answers in JSON files across restarts, takes in your edits
|
||||||
|
of those files, such as a ban you make, keep or lift, and of the rule files
|
||||||
|
while it runs, writes a JSON log line for every request, sends its log lines to
|
||||||
|
a syslog server too if you name one, serves Prometheus metrics to a scraper that
|
||||||
|
holds the metrics token, and in `observe` mode passes on the requests it would
|
||||||
|
refuse, logging what it would have done with them. 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
|
## Getting started
|
||||||
|
|
||||||
@@ -43,7 +58,9 @@ works.
|
|||||||
|
|
||||||
To work on the code, `make build` builds the binary alone, with Go installed,
|
To work on the code, `make build` builds the binary alone, with Go installed,
|
||||||
and `make run` builds and runs it, listening on port 8080 in front of an app at
|
and `make run` builds and runs it, listening on port 8080 in front of an app at
|
||||||
`SWWAF_UPSTREAM_URL`, by default `http://127.0.0.1:8081`.
|
`SWWAF_UPSTREAM_URL`, by default `http://127.0.0.1:8081`, with its state files
|
||||||
|
in `bin/state` unless `SWWAF_STATE_DIR` is set, and the default rule file of
|
||||||
|
`share/rules.d` unless `SWWAF_RULES_DIR` is set.
|
||||||
|
|
||||||
## What it does so far
|
## What it does so far
|
||||||
|
|
||||||
@@ -57,62 +74,145 @@ and `make run` builds and runs it, listening on port 8080 in front of an app at
|
|||||||
address outside `SWWAF_TRUSTED_PROXIES` is the client; if every address in it
|
address outside `SWWAF_TRUSTED_PROXIES` is the client; if every address in it
|
||||||
is inside, the leftmost is, and with no header the peer is. The app sees what
|
is inside, the leftmost is, and with no header the peer is. The app sees what
|
||||||
it would see from traefik directly: the same `Host`, the same
|
it would see from traefik directly: the same `Host`, the same
|
||||||
`X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end.
|
`X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end. It
|
||||||
- Enforces the four timeouts and the two size limits below. A limit passed
|
also gets the request's id in `X-Request-ID`, the same id as in the request's
|
||||||
before the response has started gets `smallwebwaf`'s own answer: `408` for a
|
log line (see `request_id` in "Request log" below).
|
||||||
client too slow to send its request, `413` for a request body that is too
|
- Enforces the timeouts and the size limits below. A limit passed before the
|
||||||
large, `504` for an app too slow to answer, and `502` for a response that is
|
response has started gets `smallwebwaf`'s own answer: `408` for a client too
|
||||||
too large or an app that cannot be reached. A request that announces a body
|
slow to send its request, `413` for a request body that is too large, `504`
|
||||||
over the limit is refused before anything reaches the app. While a request
|
for an app too slow to answer, and `502` for a response that is too large or
|
||||||
body is still on its way, a request timeout that runs out answers `408` if
|
an app that cannot be reached. A request that announces a body over the limit
|
||||||
`smallwebwaf` was waiting for the client to send more, and `504` if it was
|
is refused before anything reaches the app. While a request body is still on
|
||||||
waiting for the app to take what it had. Once the response has started, a
|
its way, a request timeout that runs out answers `408` if `smallwebwaf` was
|
||||||
limit can only cut the connection.
|
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
|
- 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`
|
takes the client over one of the rate limits below is refused with
|
||||||
before anything reaches the app, and so is each request after it until the
|
`SWWAF_BAN_RESPONSE`, `403` by default, before anything reaches the app, and
|
||||||
client is back under every limit. A client is one IPv4 address, or one IPv6
|
bans the client. A request whose path starts with one of
|
||||||
/64, since one abuser usually holds a whole /64. Refused requests count too,
|
`SWWAF_RATE_LIMIT_EXEMPT_PATHS`, as that setting below describes, is neither
|
||||||
so a client that keeps sending too fast stays refused until it slows down.
|
counted nor refused by the rate limits; the static lists, bans, the country
|
||||||
Each window is counted in two fixed buckets, the earlier one weighted by how
|
lists and the rule files still apply to it. A client is one IPv4 address, or
|
||||||
much of it the window still covers. At most 20,000 clients are kept, the least
|
one IPv6 /64, since one abuser usually holds a whole /64. Each window is
|
||||||
recently seen dropped first, and only in memory: a restart starts every client
|
counted in two fixed buckets, the earlier one weighted by how much of it the
|
||||||
afresh.
|
window still covers. At most 20,000 clients are kept, the least recently seen
|
||||||
- Refuses a request from a country you refuse with `403`, as soon as the
|
dropped first, with their history, and a restart gives no client a fresh
|
||||||
client's country is known and before its body is read; such a request is not
|
allowance (see "State files" below).
|
||||||
counted for the rate limits. While one of the country lists below is set, each
|
- Bans a client that breaks a rate limit, as "Bans" in [`SPEC.md`](SPEC.md)
|
||||||
client's country is looked up through GeoJS (see "Country and AS number
|
describes: the first ban lasts an hour, and a limit broken again within a day
|
||||||
lookup" below); with neither set, no visitor's address leaves the host. A
|
of a ban ending bans for three times as long as that ban, so 1, 3, 9, 27 and
|
||||||
client on a private, loopback or link-local address has no country and is
|
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, for a broken limit, for a clear sign of attack and by an admin. At
|
||||||
|
most `SWWAF_MAX_BANS` bans `smallwebwaf` made are kept, past, active and
|
||||||
|
permanent; past that, the earliest such ban of the netblock that has gone
|
||||||
|
longest without a request is dropped first. The bans whose cause is `admin`,
|
||||||
|
those you make or keep, are kept besides, and never dropped. `bans.json` shows
|
||||||
|
the bans and their notes, a restart lifts none, and you make, keep or lift a
|
||||||
|
ban by editing it (see "State files" below).
|
||||||
|
- Checks each request against the rules of the rule files (see "Rule files"
|
||||||
|
below) after the rate limits, and before its body is read. A `log` rule that
|
||||||
|
matches is noted in the log line; a `block` rule refuses the request with
|
||||||
|
`403`, and bans no one; a `ban` rule refuses it with `SWWAF_BAN_RESPONSE` and
|
||||||
|
bans the client's netblock for a clear sign of attack. Matching stops at the
|
||||||
|
first rule that refuses. A client in `SWWAF_ALLOW_NETS` is not checked.
|
||||||
|
- Bans a client for a clear sign of attack, as "Bans" in [`SPEC.md`](SPEC.md)
|
||||||
|
describes: the first such ban lasts `SWWAF_ATTACK_BAN_DURATION`, seven days by
|
||||||
|
default, and any request from the netblock while it lasts makes it permanent.
|
||||||
|
Once it has run out, the netblock is served like any other, but its next clear
|
||||||
|
sign of attack bans it permanently at once. Such a ban covers the same
|
||||||
|
netblock as a ban for a broken rate limit, does not set the client's counters
|
||||||
|
back to zero, and does not make the netblock's next ban for a broken limit
|
||||||
|
longer. Its notes give the id and the target of the rule that matched in place
|
||||||
|
of the limit.
|
||||||
|
- 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
|
never looked up: `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses it unless it is
|
||||||
in `SWWAF_ALLOW_NETS`, and `SWWAF_DENIED_COUNTRIES` does not refuse it.
|
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
|
- Checks the client's own address against the static lists, the three netblock
|
||||||
settings below, before anything else, its country included. A client in
|
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
|
`SWWAF_ALLOW_NETS` skips bans, the country lists, the rate limits and the rule
|
||||||
looked up; the timeouts and size limits still apply. A client in
|
files, and is not looked up; the timeouts and size limits still apply. A
|
||||||
`SWWAF_DENY_NETS` is refused with `403` before its body is read, and the
|
client in `SWWAF_DENY_NETS` is refused with `SWWAF_BAN_RESPONSE` before its
|
||||||
request is not counted for the rate limits; an address in `SWWAF_ALLOW_NETS`
|
body is read, and the request is not counted for the rate limits; an address
|
||||||
too is let through. A client in `SWWAF_RATE_LIMIT_EXEMPT_NETS` is neither
|
in `SWWAF_ALLOW_NETS` too is let through. A client in
|
||||||
counted nor refused by the rate limits; the country lists still apply to it.
|
`SWWAF_RATE_LIMIT_EXEMPT_NETS` is neither counted nor refused by the rate
|
||||||
|
limits; the country lists, the rule files and bans still apply to it.
|
||||||
|
- In `observe` mode, with `SWWAF_MODE=observe`, refuses none of the requests
|
||||||
|
that `SWWAF_DENY_NETS`, a ban, the country lists, a rate limit or a rule would
|
||||||
|
refuse: it passes them to the app, and their log lines name what `enforce`
|
||||||
|
mode would have done (see `would_action` in "Request log" below). The checks
|
||||||
|
run, and requests are counted, as in `enforce` mode, with three differences:
|
||||||
|
neither a broken rate limit nor a `ban` rule makes a ban; a broken rate limit
|
||||||
|
does not set the client's counters back to zero, so each request over the
|
||||||
|
limit is logged as one that would be refused; and a request under a ban does
|
||||||
|
not make it permanent. The bans in `bans.json` are kept, and refuse requests
|
||||||
|
again when `smallwebwaf` next runs in `enforce` mode, as long as they last.
|
||||||
|
The timeouts and size limits still apply, since they protect `smallwebwaf` and
|
||||||
|
the app themselves, and a request for the metrics without the token is still
|
||||||
|
answered `401`. It is for trying a configuration before enforcing it.
|
||||||
- Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any
|
- Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any
|
||||||
check and without asking the app, for the image's health check.
|
check and without asking the app, for the image's health check.
|
||||||
|
- Answers `GET /_smallwebwaf/metrics` with its metrics (see "Metrics" below) for
|
||||||
|
a request that carries `SWWAF_METRICS_TOKEN` as
|
||||||
|
`Authorization: Bearer <token>`, and with `401` for one that does not. While
|
||||||
|
the token is unset the metrics answer `404`, as does any other request under
|
||||||
|
`/_smallwebwaf/`. Unlike the health check, such a request goes through every
|
||||||
|
check any other request goes through, and is answered where another would be
|
||||||
|
passed to the app: a banned client stays refused, and each counts toward the
|
||||||
|
client's rate limits. None of them reaches the app.
|
||||||
- Writes a line in the request log for each request (see "Request log" below).
|
- Writes a line in the request log for each request (see "Request log" below).
|
||||||
|
- Sends every line it writes on stdout to a syslog server as well, while
|
||||||
|
`SWWAF_LOG_REMOTE_URL` names one (see "Sending the log to a syslog server"
|
||||||
|
below).
|
||||||
|
|
||||||
## Settings
|
## Settings
|
||||||
|
|
||||||
Each setting is an environment variable, and each has a default, so none has to
|
Each setting is an environment variable, or a file one names (see "Settings
|
||||||
be set. A setting that is set but invalid stops the start with a message naming
|
given as files" below), and each has a default, so none has to be set. A setting
|
||||||
it, and the effective settings are logged at start.
|
that is set but invalid stops the start with a message naming it, and the
|
||||||
|
effective settings are logged at start.
|
||||||
|
|
||||||
- `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens.
|
- `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens.
|
||||||
- `SWWAF_UPSTREAM_URL` (default `http://127.0.0.1:8081`): the app, as `http` or
|
- `SWWAF_UPSTREAM_URL` (default `http://127.0.0.1:8081`): the app, as `http` or
|
||||||
`https`, a host and an optional port, and nothing more.
|
`https`, a host and an optional port, and nothing more.
|
||||||
|
- `SWWAF_INSTANCE_NAME` (default: the host's name, which docker sets to the
|
||||||
|
first 12 characters of the container's id unless the deployment names one):
|
||||||
|
the name each request log line gives as `instance`. Set it, for example to
|
||||||
|
`fsn1app1/gitea`, for a name that stays the same when a deploy replaces the
|
||||||
|
container, and that tells instances apart when several log to one place.
|
||||||
|
- `SWWAF_MODE` (default `enforce`): `enforce`, or `observe` to pass on the
|
||||||
|
requests `smallwebwaf` would refuse and log what it would have done (see "What
|
||||||
|
it does so far" above).
|
||||||
- `SWWAF_TRUSTED_PROXIES` (default `10.0.0.0/8,172.16.0.0/12,192.168.0.0/16`,
|
- `SWWAF_TRUSTED_PROXIES` (default `10.0.0.0/8,172.16.0.0/12,192.168.0.0/16`,
|
||||||
the private address ranges): the netblocks whose `X-Forwarded-For` is
|
the private address ranges): the netblocks whose `X-Forwarded-For` is
|
||||||
believed. A list given replaces the default; set but empty, it trusts nothing.
|
believed. A list given replaces the default; set but empty, it trusts nothing.
|
||||||
- `SWWAF_CLIENT_REQUEST_TIMEOUT` (default `60s`): how long a client may take to
|
- `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
|
send its request line and headers, and then, from the end of the headers, its
|
||||||
body.
|
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
|
- `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.
|
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
|
- `SWWAF_UPSTREAM_REQUEST_TIMEOUT` (default `60s`): how long connecting to the
|
||||||
@@ -121,8 +221,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.
|
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_REQUEST_MAX_BYTES` (default `100M`): the largest request body.
|
||||||
- `SWWAF_RESPONSE_MAX_BYTES` (default `5G`): the largest response body.
|
- `SWWAF_RESPONSE_MAX_BYTES` (default `5G`): the largest response body.
|
||||||
- `SWWAF_ALLOW_NETS` (default empty): netblocks whose clients skip the country
|
- `SWWAF_ALLOW_NETS` (default empty): netblocks whose clients skip bans, the
|
||||||
lists and the rate limits, such as your monitoring or your own networks.
|
country lists, the rate limits and the rule files, such as your monitoring or
|
||||||
|
your own networks.
|
||||||
- `SWWAF_RATE_LIMIT_EXEMPT_NETS` (default empty): netblocks whose clients the
|
- `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.
|
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.
|
- `SWWAF_DENY_NETS` (default empty): netblocks whose clients are always refused.
|
||||||
@@ -131,12 +232,88 @@ it, and the effective settings are logged at start.
|
|||||||
requests a client may make in a minute, an hour and a day. The defaults are
|
requests a client may make in a minute, an hour and a day. The defaults are
|
||||||
several times what one busy person produces, since a browser loading a heavy
|
several times what one busy person produces, since a browser loading a heavy
|
||||||
page makes a few hundred requests and several people often share one address.
|
page makes a few hundred requests and several people often share one address.
|
||||||
|
- `SWWAF_RATE_LIMIT_EXEMPT_PATHS` (default empty): path prefixes whose requests
|
||||||
|
the rate limits neither count nor refuse, such as `/assets/` for static
|
||||||
|
assets; each starts with `/`. A request whose path, percent-decoded, contains
|
||||||
|
`..` anywhere or a backslash, or whose path as sent holds an encoded slash
|
||||||
|
(`%2F` or `%2f`), is never exempt, since the app may act on it as a path
|
||||||
|
outside every prefix: `/assets/..%2Flogin` as `/login`. Any other request is
|
||||||
|
exempt when its path as sent, the path the app receives, before any query
|
||||||
|
string and not percent-decoded, starts with a prefix, character for character.
|
||||||
|
`/assets/` matches `/assets/app.js` and `/assets/`, but not `/assets`,
|
||||||
|
`/Assets/app.js`, `/%61ssets/app.js`, `/static/assets/app.js`,
|
||||||
|
`/static/../assets/app.js` or `/assets%2Fapp.js`. A character the client sends
|
||||||
|
percent-encoded, such as a space, is written percent-encoded in a prefix, as
|
||||||
|
in `/my%20files/`, and there are no wildcards: `*` is a character like any
|
||||||
|
other.
|
||||||
- `SWWAF_DENIED_COUNTRIES` (default empty): countries whose clients are refused,
|
- `SWWAF_DENIED_COUNTRIES` (default empty): countries whose clients are refused,
|
||||||
for example `cn,ru,kp`.
|
for example `cn,ru,kp`.
|
||||||
- `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` (default empty): when set, the only
|
- `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` (default empty): when set, the only
|
||||||
countries whose clients get through, for example `us,de`. A client whose
|
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
|
country cannot be found is refused too, so that new clients are not let in
|
||||||
whenever GeoJS stops answering.
|
whenever GeoJS stops answering.
|
||||||
|
- `SWWAF_BAN_RESPONSE` (default `403`): how a refused client is answered, one
|
||||||
|
that is banned, breaks a rate limit, matches a `ban` rule, 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. A `block` rule always answers `403`.
|
||||||
|
- `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, other than one for a clear sign of attack,
|
||||||
|
bans for three times as long as that ban.
|
||||||
|
- `SWWAF_MAX_BAN_DURATION` (default `7d`): a ban for a broken rate limit that
|
||||||
|
would be longer is permanent instead.
|
||||||
|
- `SWWAF_ATTACK_BAN_DURATION` (default `7d`): the ban for a first clear sign of
|
||||||
|
attack.
|
||||||
|
- `SWWAF_MAX_BANS` (default `5000`): the most bans `smallwebwaf` made that are
|
||||||
|
kept, past, active and permanent. The bans you make or keep are kept besides.
|
||||||
|
- `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.
|
||||||
|
- `SWWAF_LOG_REQUEST_HEADERS` (default
|
||||||
|
`accept,accept-language,accept-encoding,content-type,origin,range`): the
|
||||||
|
request headers whose values the request log gives, in either case.
|
||||||
|
`Authorization`, `Cookie` and `Set-Cookie` are never logged, even when listed
|
||||||
|
(see "Request log" below). An entry naming `Host` or `Transfer-Encoding` stops
|
||||||
|
the start, since Go's HTTP server takes both out of the request; the request's
|
||||||
|
host is the field `host`.
|
||||||
|
- `SWWAF_METRICS_TOKEN` (default unset): the token a scraper sends for the
|
||||||
|
metrics, a long random value. While it is unset the metrics are off; one
|
||||||
|
shorter than 32 characters stops the start. The settings logged at start show
|
||||||
|
`********` in its place. Given as a file, it can be kept out of the app's
|
||||||
|
reach (see "Settings given as files" below).
|
||||||
|
- `SWWAF_METRICS_TOP_N` (default `50`): how many countries get series of their
|
||||||
|
own in the metrics by country; the others are counted as `other`.
|
||||||
|
- `SWWAF_RULES_DIR` (default `/etc/smallwebwaf/rules.d`): the directory of the
|
||||||
|
rule files. A directory that does not exist stops the start.
|
||||||
|
- `SWWAF_RULES_ENABLED` (default `true`): `false` reads no rule file, and checks
|
||||||
|
no request against one.
|
||||||
|
- `SWWAF_LOG_REMOTE_URL` (default unset): a syslog server that every line on
|
||||||
|
stdout is also sent to, as `syslog+udp://`, `syslog+tcp://` or `syslog+tls://`
|
||||||
|
with a host and a port, such as `syslog+tls://logs.example:6514`. Unset or
|
||||||
|
empty, nothing is sent.
|
||||||
|
- `SWWAF_LOG_REMOTE_TLS_CA_FILE` (default unset): a file of PEM certificates,
|
||||||
|
which the certificate of a `syslog+tls` server must chain to instead of the
|
||||||
|
host's own. A file that cannot be read or holds no certificate stops the
|
||||||
|
start.
|
||||||
|
- `SWWAF_LOG_REMOTE_BUFFER` (default `10000`): the most lines held while they
|
||||||
|
wait to be sent.
|
||||||
|
- `SWWAF_LOG_REMOTE_FACILITY` (default `local0`): the syslog facility the lines
|
||||||
|
are sent with: `kern`, `user`, `mail`, `daemon`, `auth`, `syslog`, `lpr`,
|
||||||
|
`news`, `uucp`, `cron`, `authpriv`, `ftp`, or `local0` to `local7`.
|
||||||
|
- `SWWAF_LOG_REMOTE_APP_NAME` (default `SWWAF_INSTANCE_NAME`): the app name the
|
||||||
|
lines are sent with, 1 to 48 printable ASCII characters without a space. While
|
||||||
|
`SWWAF_LOG_REMOTE_URL` is set, an `SWWAF_INSTANCE_NAME` that is not such a
|
||||||
|
name stops the start too, unless this setting gives one that is.
|
||||||
|
|
||||||
Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are
|
Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are
|
||||||
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
|
bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024
|
||||||
@@ -145,16 +322,46 @@ 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
|
ISO 3166-1 assigns today, and `xk` for Kosovo, in either case (`de` and `DE` are
|
||||||
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
|
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
|
||||||
`su`, stops the start, and so does a code on both country lists. `off` switches
|
`su`, stops the start, and so does a code on both country lists. `off` switches
|
||||||
a timeout, a size limit or a rate limit off.
|
a timeout, a size limit or a rate limit off;
|
||||||
|
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings, the state settings,
|
||||||
|
`SWWAF_METRICS_TOP_N` and `SWWAF_LOG_REMOTE_BUFFER` cannot be off.
|
||||||
|
|
||||||
Several limits are fixed rather than settings. The request line and headers may
|
Several limits are fixed rather than settings. At most 20,000 clients are kept,
|
||||||
take up to 32 KiB, above which the answer is `431` and nothing reaches the app.
|
with their counters and history, and an IPv6 client is counted by its /64. A new
|
||||||
A kept-open connection that sends nothing for 120 seconds is closed. That is
|
client waits at most a second for its country, and at most 100,000 answers from
|
||||||
longer than the 90 seconds after which traefik closes a connection it is not
|
GeoJS are kept, for 7 days each.
|
||||||
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
|
### Settings given as files
|
||||||
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.
|
Any setting may instead be given as a file that holds its value: the variable
|
||||||
|
named like the setting with `_FILE` added, such as `SWWAF_METRICS_TOKEN_FILE`,
|
||||||
|
names the file. `smallwebwaf` reads the file once, at start, and its health
|
||||||
|
check reads only the files of `SWWAF_LISTEN_ADDR` and `SWWAF_UPSTREAM_URL`, each
|
||||||
|
time it runs. The file's contents are the value, less one newline at their end
|
||||||
|
so that a file written with `echo` or an editor works, and are checked as the
|
||||||
|
setting's own value would be. Setting both the setting and its `_FILE` form, or
|
||||||
|
naming a file that cannot be read, stops the start with a message naming the
|
||||||
|
variable. The settings logged at start name the file, and show a token given in
|
||||||
|
one as `********`, as they show one given directly.
|
||||||
|
`SWWAF_LOG_REMOTE_TLS_CA_FILE`, whose value names a file already, has no `_FILE`
|
||||||
|
form.
|
||||||
|
|
||||||
|
The app starts with the same environment variables as `smallwebwaf`, so it can
|
||||||
|
read a token given as one. A token given as a file is out of the app's reach
|
||||||
|
only while the `smallwebwaf` user alone can read the file: make it on the host,
|
||||||
|
owned by uid 65532, the `smallwebwaf` user, with mode `0400`, and mount the
|
||||||
|
directory that holds it into the container read-only; the container sees the
|
||||||
|
same owner and mode. For example, on the host:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
mkdir -p /srv/app/tokens
|
||||||
|
openssl rand -hex 32 > /srv/app/tokens/metrics
|
||||||
|
chown 65532:65532 /srv/app/tokens/metrics
|
||||||
|
chmod 0400 /srv/app/tokens/metrics
|
||||||
|
```
|
||||||
|
|
||||||
|
and for the container, `-v /srv/app/tokens:/etc/smallwebwaf/tokens:ro` and
|
||||||
|
`-e SWWAF_METRICS_TOKEN_FILE=/etc/smallwebwaf/tokens/metrics`.
|
||||||
|
|
||||||
## Request log
|
## Request log
|
||||||
|
|
||||||
@@ -162,41 +369,352 @@ at most 100,000 answers from GeoJS are kept, for 7 days each.
|
|||||||
refused ones included:
|
refused ones included:
|
||||||
|
|
||||||
```
|
```
|
||||||
{"type":"request","time":"2026-10-03T12:00:00.123Z","client_ip":"203.0.113.9","peer_ip":"172.18.0.2","country":"DE","method":"GET","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"upstream_status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","action":"forward","duration_total":3.217,"duration_upstream_total":3.104}
|
{"type":"request","time":"2026-10-03T12:00:00.123Z","instance":"fsn1app1/gitea","client_ip":"203.0.113.9","method":"GET","scheme":"https","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","request_id":"7Q2NHZ4KJ3VXW5YB6R3MEFTD2A","peer_ip":"172.18.0.2","forwarded_for":"203.0.113.9","client_group":"203.0.113.9/32","country":"DE","request_headers":{"accept":"*/*"},"response_content_type":"text/html; charset=utf-8","upstream_status":200,"action":"forward","counts":{"minute":1,"hour":12,"day":40},"duration_total":3.217,"duration_checks":0.041,"duration_upstream_connect":0.052,"duration_upstream_first_byte":2.874,"duration_upstream_total":3.104}
|
||||||
```
|
```
|
||||||
|
|
||||||
- `time` is when the request arrived, in UTC. `peer_ip` is the TCP peer,
|
A field that does not apply to a request is left out of its line, apart from
|
||||||
normally traefik. `path` and `query` are as the client sent them.
|
`type`, the fields from `time` to `user_agent`, `request_id`, `peer_ip`,
|
||||||
- `country` is the client's country as GeoJS places it, and empty when it is not
|
`client_group`, `country`, `action` and `duration_total`, which every line has.
|
||||||
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,
|
- `time` is when the request arrived, in UTC. `instance` is
|
||||||
and when GeoJS cannot place the client or has not answered in time.
|
`SWWAF_INSTANCE_NAME`. `scheme` is the `X-Forwarded-Proto` a trusted proxy
|
||||||
|
sent, and otherwise `http`. `path` and `query` are as the client sent them.
|
||||||
|
- `request_id` is the `X-Request-ID` a trusted proxy sent, or a new random one
|
||||||
|
of 26 letters and digits when it sent none, or when the peer is not a trusted
|
||||||
|
proxy. A request passed to the app takes it there in `X-Request-ID`.
|
||||||
|
- `peer_ip` is the TCP peer, normally traefik. `forwarded_for` is the
|
||||||
|
`X-Forwarded-For` header as received, several lines of it joined with `, `.
|
||||||
|
`client_group` is the client as the rate limits count it: its IPv4 address as
|
||||||
|
a /32, or the /64 of its IPv6 address.
|
||||||
|
- `country` is the client's country as GeoJS places it. It is empty with neither
|
||||||
|
country 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 whose client a ban
|
||||||
|
covers, even when the client's country is known.
|
||||||
|
- `content_type` is the request's `Content-Type`, and `content_length` the
|
||||||
|
length the request announced for its body, which is left out for none or zero.
|
||||||
|
- `request_headers` are the request's headers that `SWWAF_LOG_REQUEST_HEADERS`
|
||||||
|
names, by name in lower case, several lines of one joined with `, `.
|
||||||
|
`Authorization`, `Cookie` and `Set-Cookie` are never among them, whatever the
|
||||||
|
setting says: `has_authorization` and `has_cookie` are there instead, and
|
||||||
|
true, when the request has an `Authorization` or a `Cookie` header.
|
||||||
|
- `websocket` is there, and true, when the app switched the connection to
|
||||||
|
another protocol, as it does for a WebSocket.
|
||||||
- `status` is what the client was sent, `0` if nothing was; `upstream_status` is
|
- `status` is what the client was sent, `0` if nothing was; `upstream_status` is
|
||||||
what the app answered, and is left out when the app did not answer.
|
what the app answered, and is left out when the app did not answer.
|
||||||
|
- `response_content_type`, `cache_control` and `location` are the
|
||||||
|
`Content-Type`, `Cache-Control` and `Location` headers of the answer: the
|
||||||
|
app's, as passed on, or those of `smallwebwaf`'s own answer.
|
||||||
- `request_bytes` and `response_bytes` count body bytes.
|
- `request_bytes` and `response_bytes` count body bytes.
|
||||||
- `action` is `forward` for a request passed to the app, `denied` for one
|
- `action` is `forward` for a request passed to the app, `denied` for one
|
||||||
refused because its client is in `SWWAF_DENY_NETS`, `country_denied` for one
|
refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused
|
||||||
refused for its client's country, `rate_limited` for one refused for a rate
|
because a ban covers its client or because it matched a `ban` rule, which bans
|
||||||
limit, `too_large` for a request or response over its size limit, `timed_out`
|
its client, `country_denied` for one refused for its client's country,
|
||||||
for one that ran out of time, `upstream_error` when the app could not be
|
`rate_limited` for one that broke a rate limit and banned its client,
|
||||||
reached or its answer broke off, and `admin` for one `smallwebwaf` answered at
|
`rule_blocked` for one a `block` rule refused, `too_large` for a request or
|
||||||
its own endpoint.
|
response over its size limit, `timed_out` for one that ran out of time,
|
||||||
- `limit_hit` is there for a request refused for a rate limit, and names the
|
`upstream_error` when the app could not be reached or its answer broke off,
|
||||||
|
and `admin` for one `smallwebwaf` answered at its own endpoint.
|
||||||
|
- `would_action` is there in `observe` mode for a request that
|
||||||
|
`SWWAF_DENY_NETS`, a ban, the country lists, a rate limit or a rule would have
|
||||||
|
refused in `enforce` mode, and names the action that refusal would have had:
|
||||||
|
`denied`, `banned`, `country_denied`, `rate_limited` or `rule_blocked`.
|
||||||
|
`action` then names what was done: `forward` for a request passed to the app,
|
||||||
|
and another action, such as `too_large`, for one a size or time limit refused.
|
||||||
|
- `counts` gives the client's requests in the minute, the hour and the day as
|
||||||
|
the rate limits count them, this request included: in each window, those in
|
||||||
|
the bucket under way and a share of those in the bucket before, so a count can
|
||||||
|
have a fraction. For a request that broke a limit, they are the counts that
|
||||||
|
broke it. It is left out for a request the rate limits do not count: the
|
||||||
|
health check, one from a client in `SWWAF_ALLOW_NETS` or
|
||||||
|
`SWWAF_RATE_LIMIT_EXEMPT_NETS`, one for a path that
|
||||||
|
`SWWAF_RATE_LIMIT_EXEMPT_PATHS` exempts, and one that `SWWAF_DENY_NETS`, a ban
|
||||||
|
or the country lists refuse, or would refuse in `observe` mode. The byte
|
||||||
|
totals come with the byte limits.
|
||||||
|
- `rule_ids` is there for a request that matched rules of the rule files, and
|
||||||
|
lists their ids in the order they matched, up to the one that refused it.
|
||||||
|
- `limit_hit` is there for a request that broke a rate limit, and names the
|
||||||
window whose limit it went over: `minute`, `hour` or `day`, the shortest if it
|
window whose limit it went over: `minute`, `hour` or `day`, the shortest if it
|
||||||
went over several.
|
went over several. `offence` is then `limit`.
|
||||||
|
- `ban_expires` is there for a request that made a ban or was refused under one,
|
||||||
|
or in `observe` mode would have been 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.
|
- `aborted` is there, and true, when the client went away early.
|
||||||
- `duration_total` and `duration_upstream_total` are in milliseconds.
|
- The timings are in milliseconds, to the microsecond. `duration_total` runs
|
||||||
|
from when the request's headers had been read to when its line is written, and
|
||||||
|
`duration_checks` over the same start to when the checks were done; the health
|
||||||
|
check runs none, and its line has no `duration_checks`.
|
||||||
|
`duration_upstream_connect`, `duration_upstream_first_byte` and
|
||||||
|
`duration_upstream_total` are there for a request passed to the app, and run
|
||||||
|
from when it was handed to the app: until there was a connection to it, new or
|
||||||
|
kept open from an earlier request, until the first byte of its answer arrived,
|
||||||
|
and until the end. The first two are left out when that never happened, as for
|
||||||
|
an app that cannot be reached.
|
||||||
|
|
||||||
No body and no other header is logged. `smallwebwaf`'s own messages (start, the
|
No body is logged, and no header but those above. `smallwebwaf`'s own messages
|
||||||
settings, stop, errors) share the stream as JSON lines marked
|
(start, the settings, stop, errors) share the stream as JSON lines marked
|
||||||
`"type":"process"`.
|
`"type":"process"`.
|
||||||
|
|
||||||
Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and
|
Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and
|
||||||
headers before `smallwebwaf` sees the request, and some requests end there,
|
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
|
without a line in the log: headers over `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`,
|
||||||
slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`, whose connection it closes without
|
which it answers `431`, headers slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`,
|
||||||
an answer, and requests it cannot read at all, which it answers itself, mostly
|
whose connection it closes without an answer, and requests it cannot read at
|
||||||
with `400`.
|
all, which it answers itself, mostly with `400`.
|
||||||
|
|
||||||
|
### Sending the log to a syslog server
|
||||||
|
|
||||||
|
While `SWWAF_LOG_REMOTE_URL` is set, every line `smallwebwaf` writes on stdout,
|
||||||
|
request lines and its own, is also sent to that syslog server, as the message of
|
||||||
|
an RFC 5424 record: one record to a datagram over UDP, and over TCP and TLS each
|
||||||
|
record after its length in bytes and a space. A record gives the facility
|
||||||
|
`SWWAF_LOG_REMOTE_FACILITY` names, the severity informational, the time the line
|
||||||
|
was written, in the same form as a request line's `time`, the host's name, and
|
||||||
|
the app name `SWWAF_LOG_REMOTE_APP_NAME` gives. stdout is unchanged.
|
||||||
|
|
||||||
|
The lines wait in a buffer of `SWWAF_LOG_REMOTE_BUFFER` lines and are sent from
|
||||||
|
there, so a server that is slow or cannot be reached never holds up a request or
|
||||||
|
stdout. When the buffer is full, its oldest line is dropped to make room. A line
|
||||||
|
whose sending fails is dropped too, and the connection closed. That failure,
|
||||||
|
like a failed attempt to connect, is logged and followed by the next attempt to
|
||||||
|
connect a second later, twice as long after each further failure up to a minute,
|
||||||
|
and a second again after a connection that stayed up for a minute before it
|
||||||
|
failed. A line too long for one UDP datagram is dropped alone, with no wait and
|
||||||
|
nothing logged. UDP gives no sign of what arrives, and over TCP and TLS a line
|
||||||
|
sent on a connection the server has just closed can be lost before a failure
|
||||||
|
shows; such a loss is not counted.
|
||||||
|
|
||||||
|
As `smallwebwaf` stops, it sends the lines still waiting, on the connection open
|
||||||
|
or a new one, for at most two seconds, and gives up the rest; stdout has carried
|
||||||
|
them.
|
||||||
|
|
||||||
|
## 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`. A ban's `cause` is `limit` for a broken rate limit or
|
||||||
|
`attack` for a clear sign of attack, for a ban `smallwebwaf` made, and `admin`
|
||||||
|
for one you made or keep. Its `reason` is a short text: for a ban
|
||||||
|
`smallwebwaf` made, the limit broken, such as
|
||||||
|
`requests per minute over the limit of 1000`, or the rule that matched, such
|
||||||
|
as `matched the rule env-file`; for yours, what you wrote. Its `lifted` is
|
||||||
|
when you lifted it, and is left out until you do.
|
||||||
|
- `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 (one
|
||||||
|
`smallwebwaf` answered at its own endpoints is neither, unless it was refused
|
||||||
|
with `401` for a missing or wrong token), 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 or made
|
||||||
|
permanent, with every such change 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`. So does a ban whose `cause` is not `limit`, `attack` or `admin`. The
|
||||||
|
AS number and AS name come with their lookup.
|
||||||
|
|
||||||
|
While it runs, `smallwebwaf` watches `SWWAF_STATE_DIR` and takes in your edit of
|
||||||
|
a state file as soon as you save it: what the file then holds replaces what
|
||||||
|
`smallwebwaf` held for it, as if read at start. It tells its own writes from
|
||||||
|
yours by comparing the file with what it last read or wrote, and before it
|
||||||
|
writes a file it takes in any edit made since, so your edit is not overwritten;
|
||||||
|
a change `smallwebwaf` made after you opened the file, such as a new ban, is
|
||||||
|
lost when you save over it. An edit that would stop the start, because it does
|
||||||
|
not parse, has another `version`, leaves out a field an entry needs or gives a
|
||||||
|
ban another `cause`, does not stop the running `smallwebwaf`: it keeps what it
|
||||||
|
holds, and at the file's next write renames your file to `<name>.bad`, such as
|
||||||
|
`bans.json.bad`, writes the file again from memory, and logs the file and where
|
||||||
|
the error is. It waits for that write because an editor's file can be read
|
||||||
|
before the editor has finished writing it. Mend the `.bad` file and move it
|
||||||
|
back. A file you remove is written again at its next write.
|
||||||
|
|
||||||
|
To ban a netblock, add an entry to `bans.json` with its `netblock`, its `start`
|
||||||
|
and its `expires`, `null` for a ban that never ends; its `reason` and its
|
||||||
|
`notes` may be left out, and so may its `cause`, which is then `admin`, and is
|
||||||
|
written so at the file's next write. A ban whose `cause` is `admin` is never
|
||||||
|
dropped and does not count toward `SWWAF_MAX_BANS`. A ban whose `cause` is
|
||||||
|
`attack` becomes permanent at the first request it refuses; one whose `cause` is
|
||||||
|
`admin` does not. This `bans.json` bans `203.0.113.0/24` for good:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"bans": [
|
||||||
|
{
|
||||||
|
"netblock": "203.0.113.0/24",
|
||||||
|
"start": "2026-10-06T12:00:00Z",
|
||||||
|
"expires": null,
|
||||||
|
"reason": "probes for logins"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
To keep a ban `smallwebwaf` made, so that it is never dropped, set its `cause`
|
||||||
|
to `admin`: `"cause": "admin"`.
|
||||||
|
|
||||||
|
To lift a ban, add `lifted` to its entry, with the time you lift it, such as
|
||||||
|
`"lifted": "2026-10-06T13:00:00Z"`. From when the edit is taken in, the ban
|
||||||
|
refuses nothing, whatever time `lifted` gives, and does not make the netblock's
|
||||||
|
next ban longer; it is kept in `bans.json` with its notes, as any other ban is.
|
||||||
|
To forget a ban altogether, delete its entry: it then refuses nothing either,
|
||||||
|
and does not make the netblock's next ban longer.
|
||||||
|
|
||||||
|
## Rule files
|
||||||
|
|
||||||
|
`smallwebwaf` reads every `*.rules` file in `SWWAF_RULES_DIR`,
|
||||||
|
`/etc/smallwebwaf/rules.d` by default, in the order of their names, and checks
|
||||||
|
each request against their rules in that order, as "Rule files" in
|
||||||
|
[`SPEC.md`](SPEC.md) describes. A file whose name starts with `.`, such as an
|
||||||
|
editor's lock file `.#50-app.rules`, is not a rule file, as a shell's `*.rules`
|
||||||
|
would not match it. A rule is a line of four fields separated by spaces or tabs:
|
||||||
|
an id, a target, an action and a regex, which runs to the end of the line.
|
||||||
|
Spaces and tabs at the end of a line are not part of its regex, so a line with
|
||||||
|
only those after its action has no regex, and is not a rule. Blank lines and
|
||||||
|
lines that start with `#` are ignored.
|
||||||
|
|
||||||
|
```
|
||||||
|
# id target action regex
|
||||||
|
env-file path ban (?i)^/\.env(\.[a-z]+)?$
|
||||||
|
scanner-agent user_agent ban (?i)\b(sqlmap|nikto|nuclei|wpscan)\b
|
||||||
|
```
|
||||||
|
|
||||||
|
- The id is letters, digits, `-` and `_`, and no two rules share one. The
|
||||||
|
request log, the metrics and a ban's notes name the rule by it.
|
||||||
|
- The target is what the regex is matched against: `path` or `query`, as the
|
||||||
|
client sent it, before any decoding; `uri`, the path and the query together,
|
||||||
|
both as sent and once percent-decoded, so that an encoded probe does not slip
|
||||||
|
past; `method`; `host`; `user_agent`; `referer`; or `header:<Name>`, any one
|
||||||
|
request header but `Host` and `Transfer-Encoding`, which Go's HTTP server
|
||||||
|
takes out of every request; the request's host is the target `host`. A header
|
||||||
|
sent more than once is matched with its values joined by `, `, and one not
|
||||||
|
sent as empty text. No body is read.
|
||||||
|
- The action is `log`, `block` or `ban` (see "What it does so far" above). Keep
|
||||||
|
`ban` for requests no real visitor sends, and anchor a path at the site root
|
||||||
|
with `^/`: a file of the same name deeper in a site can be ordinary content,
|
||||||
|
such as a file in a repository on a code forge.
|
||||||
|
- The regex is in Go's syntax, RE2, which has no backreferences or lookaround,
|
||||||
|
and takes time linear in the text it reads. It matches anywhere in the target
|
||||||
|
unless anchored with `^` and `$`; `(?i)` at its front makes it ignore case.
|
||||||
|
|
||||||
|
A line that is not a rule, a `header:` name with a character no header name can
|
||||||
|
have, such as `header:User-Agent:`, a rule for the `Host` or the
|
||||||
|
`Transfer-Encoding` header, a regex that does not compile or an id used twice
|
||||||
|
stops the start with a message naming the file and the line, and so does a
|
||||||
|
`SWWAF_RULES_DIR` that does not exist. An empty directory is no error, and the
|
||||||
|
log says that it holds no rules. While it runs, `smallwebwaf` watches the
|
||||||
|
directory, and reads the rule files again once the directory has had no change
|
||||||
|
for 2 seconds after one is edited, added or removed, so that a file saved in
|
||||||
|
place, appended to or copied in with `scp` is read only once whole, unless its
|
||||||
|
writing stops for longer. It also reads them 2 seconds after it starts watching,
|
||||||
|
so that an edit saved while it started is not missed. If they then hold one of
|
||||||
|
those errors, the rules stay as they were, the earlier version of the edited
|
||||||
|
file included, the log names the file and the line, and the files are read again
|
||||||
|
after the next change.
|
||||||
|
|
||||||
|
The image ships one rule file, `share/rules.d/00-default.rules` here: rules that
|
||||||
|
ban probes no real visitor sends, for secrets, version control directories,
|
||||||
|
backups, logs and web shells at the site root, and the user agents of common
|
||||||
|
scanners; one that blocks `../` twice in a row in the path or the query; and one
|
||||||
|
that only notes a request without a user agent. An app's Dockerfile adds rules
|
||||||
|
of its own in a file beside it, named to sort after it, such as this
|
||||||
|
`50-gitea.rules` for an app that serves no WordPress:
|
||||||
|
|
||||||
|
```
|
||||||
|
wp-probe path ban (?i)^/(wp-login\.php|xmlrpc\.php|wp-admin/)
|
||||||
|
```
|
||||||
|
|
||||||
|
```dockerfile
|
||||||
|
COPY 50-gitea.rules /etc/smallwebwaf/rules.d/50-gitea.rules
|
||||||
|
```
|
||||||
|
|
||||||
|
A directory mounted over `/etc/smallwebwaf/rules.d` replaces the default file,
|
||||||
|
and single files mounted into it add to it. Docker does not show a single
|
||||||
|
mounted file being replaced, which is how many editors save, so rules to be
|
||||||
|
edited while `smallwebwaf` runs belong in a mounted directory, with a copy of
|
||||||
|
`00-default.rules` if its rules are to stay. To run without rules, mount an
|
||||||
|
empty directory or set `SWWAF_RULES_ENABLED=false`.
|
||||||
|
|
||||||
|
## Metrics
|
||||||
|
|
||||||
|
`GET /_smallwebwaf/metrics` answers with the metrics in the Prometheus text
|
||||||
|
format, for a scraper that sends `SWWAF_METRICS_TOKEN`, through traefik like any
|
||||||
|
other request. No metric carries a client's address.
|
||||||
|
|
||||||
|
- `smallwebwaf_requests_total`, `smallwebwaf_request_bytes_total` and
|
||||||
|
`smallwebwaf_response_bytes_total`: requests, and their body bytes each way,
|
||||||
|
by `status_class`, such as `2xx`, or `none` when nothing was sent, and by
|
||||||
|
`action`, as the request log names it.
|
||||||
|
- `smallwebwaf_request_duration_seconds`: how long requests took, and
|
||||||
|
`smallwebwaf_upstream_duration_seconds`: how long those passed to the app took
|
||||||
|
from then on, as histograms; `smallwebwaf_requests_in_flight`: the requests
|
||||||
|
under way.
|
||||||
|
- `smallwebwaf_rate_limit_hits_total` by `window`,
|
||||||
|
`smallwebwaf_size_and_time_limit_hits_total` by `limit`, the setting whose
|
||||||
|
limit was passed, `smallwebwaf_offences_total` by `kind`, and
|
||||||
|
`smallwebwaf_bans_made_total` by `cause`, `limit`, `attack` or `admin`, the
|
||||||
|
last for the bans whose `cause` is `admin` that you add to `bans.json` while
|
||||||
|
`smallwebwaf` runs; `smallwebwaf_active_bans` and
|
||||||
|
`smallwebwaf_permanent_bans`, neither of which counts a lifted ban.
|
||||||
|
- `smallwebwaf_rule_matches_total`: the requests that matched each rule, by
|
||||||
|
`rule_id` and `action`, the rule's own; and `smallwebwaf_rules_loaded`: the
|
||||||
|
rules read from the rule files.
|
||||||
|
- `smallwebwaf_country_requests_total`,
|
||||||
|
`smallwebwaf_country_request_bytes_total`,
|
||||||
|
`smallwebwaf_country_response_bytes_total`, and
|
||||||
|
`smallwebwaf_country_list_refusals_total`, the requests the country lists
|
||||||
|
refused, by `country`, for the requests whose client's country is known. The
|
||||||
|
`SWWAF_METRICS_TOP_N` countries with the most requests since the start have
|
||||||
|
series of their own, and the others are counted as `other`. A country that
|
||||||
|
drops out of them loses its series, and its later requests count as `other`;
|
||||||
|
one that comes into them gets a series that counts from then on.
|
||||||
|
- `smallwebwaf_geojs_requests_total`: the requests to GeoJS;
|
||||||
|
`smallwebwaf_geojs_failures_total`: those that failed, an answer that leaves
|
||||||
|
out an address asked about included; and `smallwebwaf_geojs_unanswered_total`:
|
||||||
|
the requests whose client counted as coming from an unknown country because
|
||||||
|
GeoJS had not answered about it in time.
|
||||||
|
- `smallwebwaf_tracked_clients`: the clients in the table of clients.
|
||||||
|
- `smallwebwaf_state_file_writes_total`,
|
||||||
|
`smallwebwaf_state_file_write_failures_total`,
|
||||||
|
`smallwebwaf_state_file_last_write_timestamp_seconds` and
|
||||||
|
`smallwebwaf_state_file_size_bytes`, by `file`; and, by `file` too,
|
||||||
|
`smallwebwaf_state_file_edits_taken_in_total`: your edits taken in, and
|
||||||
|
`smallwebwaf_state_file_edits_set_aside_total`: those renamed to `<name>.bad`
|
||||||
|
because they would stop the start.
|
||||||
|
- While `SWWAF_LOG_REMOTE_URL` is set,
|
||||||
|
`smallwebwaf_remote_log_lines_sent_total`: the lines sent to it;
|
||||||
|
`smallwebwaf_remote_log_lines_dropped_total`: those dropped, from a full
|
||||||
|
buffer or because their sending failed; and
|
||||||
|
`smallwebwaf_remote_log_buffer_depth`: those waiting in the buffer.
|
||||||
|
- Go's own `go_` metrics and the process's `process_` metrics.
|
||||||
|
|
||||||
|
The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request
|
||||||
|
log") are not counted. The metrics of the features still to come, such as the
|
||||||
|
Core Rule Set, come with them.
|
||||||
|
|
||||||
## Why
|
## Why
|
||||||
|
|
||||||
@@ -299,9 +817,9 @@ goes through the candidates one by one.
|
|||||||
readable JSON files, written regularly and at every stop, so a restart loses
|
readable JSON files, written regularly and at every stop, so a restart loses
|
||||||
nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks
|
nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks
|
||||||
up the change. Nothing is read from disk while serving a request. The files
|
up the change. Nothing is read from disk while serving a request. The files
|
||||||
come in milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md));
|
for the bans, the clients and the GeoJS answers are built, with an edit taken
|
||||||
until then the rate counters and the GeoJS answers are kept in memory only,
|
in while running (see "State files" above); the others come with their
|
||||||
and a restart loses them.
|
features.
|
||||||
- Health checks, the metrics, and listing, adding and lifting bans or asking why
|
- Health checks, the metrics, and listing, adding and lifting bans or asking why
|
||||||
a given address was refused, all on the one port every request uses: under
|
a given address was refused, all on the one port every request uses: under
|
||||||
`/_smallwebwaf/` on the app's own address, through traefik like any other
|
`/_smallwebwaf/` on the app's own address, through traefik like any other
|
||||||
@@ -397,8 +915,9 @@ main "$@"
|
|||||||
the health check on `127.0.0.1`.
|
the health check on `127.0.0.1`.
|
||||||
- `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf`. Mount a volume
|
- `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf`. Mount a volume
|
||||||
there to keep bans and client history when a deploy replaces the container;
|
there to keep bans and client history when a deploy replaces the container;
|
||||||
without one, it still starts. The state files come in milestone 3 or later;
|
without one, it still starts. At each start the `run` script of `smallwebwaf`
|
||||||
until then it writes nothing to disk and needs no volume.
|
gives that directory and every file in it to the `smallwebwaf` user, so a host
|
||||||
|
directory mounted there needs no change of owner.
|
||||||
- `docker stop` has runit stop both processes. `smallwebwaf` then stops taking
|
- `docker stop` has runit stop both processes. `smallwebwaf` then stops taking
|
||||||
requests and gives those in progress five seconds to finish.
|
requests and gives those in progress five seconds to finish.
|
||||||
|
|
||||||
@@ -418,28 +937,27 @@ the metrics, failure behaviour and the build order.
|
|||||||
|
|
||||||
So far `smallwebwaf` looks up only the country, only through GeoJS, and only
|
So far `smallwebwaf` looks up only the country, only through GeoJS, and only
|
||||||
while `SWWAF_DENIED_COUNTRIES` or `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` is set:
|
while `SWWAF_DENIED_COUNTRIES` or `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` is set:
|
||||||
then the address of every new visitor outside `SWWAF_ALLOW_NETS` and
|
then the address of every new visitor is sent to GeoJS, except a visitor in
|
||||||
`SWWAF_DENY_NETS` is sent to GeoJS, and with neither set, none is. An IPv6
|
`SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS` and one whose netblock a ban covers, and
|
||||||
visitor is asked about by the first address of its /64. A new visitor waits at
|
with neither set, none is. An IPv6 visitor is asked about by the first address
|
||||||
most a second for its answer, and without one counts as coming from an unknown
|
of its /64. A new visitor waits at most a second for its answer, and without one
|
||||||
country until the answer arrives. The addresses waiting are asked about
|
counts as coming from an unknown country until the answer arrives. The addresses
|
||||||
together, up to 200 in one request, one request at a time; at most 10,000
|
waiting are asked about together, up to 200 in one request, one request at a
|
||||||
visitors wait, and one more counts as coming from an unknown country until there
|
time; at most 10,000 visitors wait, and one more counts as coming from an
|
||||||
is room. While GeoJS fails, visitors with a kept answer are unaffected and new
|
unknown country until there is room. While GeoJS fails, visitors with a kept
|
||||||
ones count as coming from an unknown country. GeoJS is then left alone for a
|
answer are unaffected and new ones count as coming from an unknown country.
|
||||||
second, twice as long after each further failure up to five minutes, and asked
|
GeoJS is then left alone for a second, twice as long after each further failure
|
||||||
again by the next request that needs it.
|
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
|
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
|
client, for the request log, the metrics and the ban notes, and for the country
|
||||||
lists and biased limits when you set them. It works with no setup: by default it
|
lists and biased limits when you set them. It works with no setup: by default it
|
||||||
asks the free GeoJS web service, which needs no account and no file. This means
|
asks the free GeoJS web service, which needs no account and no file. This means
|
||||||
that, by default, the address of every new visitor is sent to GeoJS. Each answer
|
that, by default, the address of every new visitor is sent to GeoJS. Each answer
|
||||||
is kept in memory for seven days, and many addresses are asked about in one
|
is kept for seven days, in memory and in `lookups.json`, so that it survives a
|
||||||
request; writing the answers to disk, so that they survive a restart, comes in
|
restart, and many addresses are asked about in one request. GeoJS publishes no
|
||||||
milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md)). GeoJS
|
rate limit but may block a caller it thinks asks too much; while it is not
|
||||||
publishes no rate limit but may block a caller it thinks asks too much; while it
|
answering, new visitors count as coming from an unknown country, which
|
||||||
is not answering, new visitors count as coming from an unknown country, which
|
|
||||||
`SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses.
|
`SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses.
|
||||||
|
|
||||||
To keep your visitors' addresses on your own host, set
|
To keep your visitors' addresses on your own host, set
|
||||||
@@ -470,32 +988,55 @@ addresses are never sent to GeoJS.
|
|||||||
## How the code is laid out
|
## How the code is laid out
|
||||||
|
|
||||||
- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`.
|
- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`.
|
||||||
- `internal/smallwebwaf`: the process: it reads the settings, listens, serves
|
- `internal/smallwebwaf`: the process: it reads the settings, the rule files and
|
||||||
requests until `SIGTERM` or `SIGINT`, and stops. Run as
|
the state files, listens, serves requests until `SIGTERM` or `SIGINT`, and
|
||||||
`smallwebwaf healthcheck`, it is the image's health check instead.
|
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/config`: reads the settings, the one place they are read.
|
||||||
- `internal/proxy`: what happens to each request: it works out the client, runs
|
- `internal/proxy`: what happens to each request: it works out the client, runs
|
||||||
the checks, passes the request to the app and the answer back with the
|
the checks, passes the request to the app and the answer back with the
|
||||||
standard library's `httputil.ReverseProxy` within the timeouts and size
|
standard library's `httputil.ReverseProxy` within the timeouts and size
|
||||||
limits, and writes the request's log line. Its `check` method is where a
|
limits, and writes the request's log line. Its `check` method is where a
|
||||||
request is refused before anything reaches the app: for `SWWAF_DENY_NETS`, for
|
request is refused before anything reaches the app: for `SWWAF_DENY_NETS`, for
|
||||||
the country lists, for a rate limit, and for an announced body over the size
|
a ban, for the country lists, for a rate limit, which bans the client, for a
|
||||||
limit.
|
`block` or `ban` rule, the latter banning the client, and for an announced
|
||||||
|
body over the size limit; in `observe` mode, only for the size limit, with
|
||||||
|
what it would have refused for noted in the log line. A request under
|
||||||
|
`/_smallwebwaf/` that `check` lets through is answered by `answerAdmin`
|
||||||
|
instead of reaching the app.
|
||||||
|
- `internal/metrics`: the metrics, counted as the other parts tell it what
|
||||||
|
happened, and served in the Prometheus text format.
|
||||||
|
- `internal/bans`: the ban ledger: each netblock's bans with their notes, how
|
||||||
|
long a new ban lasts, when a ban for a clear sign of attack becomes permanent,
|
||||||
|
and which ban `smallwebwaf` made is dropped when `SWWAF_MAX_BANS` are held.
|
||||||
|
- `internal/rules`: reads the rule files at start and again as they change, and
|
||||||
|
tells which of their rules a request matches.
|
||||||
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the
|
- `internal/lookup`: looks up each client's country through GeoJS, and keeps the
|
||||||
answers.
|
answers.
|
||||||
- `internal/ratelimit`: counts each client's requests and tells when one takes
|
- `internal/ratelimit`: the table of clients: counts each client's requests,
|
||||||
it over a rate limit.
|
tells when one takes it over a rate limit, and keeps each client's history.
|
||||||
|
- `internal/state`: reads the state files at start, takes in an admin's edit of
|
||||||
|
one while running, and writes them when they are due and at the stop.
|
||||||
- `internal/requestlog`: the lines on stdout: the request log line and the
|
- `internal/requestlog`: the lines on stdout: the request log line and the
|
||||||
process's own messages.
|
process's own messages.
|
||||||
|
- `internal/remotelog`: sends the lines on stdout to `SWWAF_LOG_REMOTE_URL`,
|
||||||
|
each as a syslog record, from a buffer of its own. It is written with the
|
||||||
|
standard library alone, whose `log/syslog` writes only the older syslog
|
||||||
|
format.
|
||||||
- `Dockerfile`: the lint and test phases, then the image, whose last stage
|
- `Dockerfile`: the lint and test phases, then the image, whose last stage
|
||||||
installs Ubuntu's packages, nixpkgs, `runsvinit` and `smallwebwaf`, with
|
installs Ubuntu's packages, nixpkgs, `runsvinit` and `smallwebwaf`, with
|
||||||
`share/smallwebwaf.run` as runit's `run` script for `smallwebwaf`.
|
`share/smallwebwaf.run` as runit's `run` script for `smallwebwaf` and
|
||||||
|
`share/rules.d/00-default.rules` as its default rule file.
|
||||||
- `deploy/example-app`: an app built on the image, which `script/example-app`
|
- `deploy/example-app`: an app built on the image, which `script/example-app`
|
||||||
checks.
|
checks.
|
||||||
|
|
||||||
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
|
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
|
||||||
table of clients to 20,000 and the GeoJS answers to 100,000, dropping the least
|
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`.
|
recently seen, and the banned netblocks in the order they were last seen, from
|
||||||
|
which the ledger picks the ban to drop past `SWWAF_MAX_BANS`, and
|
||||||
|
`github.com/prometheus/client_golang` keeps the metrics and serves them, and
|
||||||
|
`github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file or a rule
|
||||||
|
file is saved. The country codes are the list in `internal/config/config.go`.
|
||||||
|
|
||||||
## Entrypoints
|
## Entrypoints
|
||||||
|
|
||||||
@@ -524,19 +1065,23 @@ so that they run in minimal containers.
|
|||||||
- `script/install-precommit`: installs that hook; `make hooks` runs it.
|
- `script/install-precommit`: installs that hook; `make hooks` runs it.
|
||||||
- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for
|
- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for
|
||||||
working on the code by hand; `make build` runs it.
|
working on the code by hand; `make build` runs it.
|
||||||
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it;
|
- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it, with
|
||||||
`make run` runs it.
|
its state files in `bin/state` unless `SWWAF_STATE_DIR` is set, and the rule
|
||||||
|
files of `share/rules.d` unless `SWWAF_RULES_DIR` is set; `make run` runs it.
|
||||||
- `script/example-app`: builds the image and, on it, the example app in
|
- `script/example-app`: builds the image and, on it, the example app in
|
||||||
`deploy/example-app`, runs it, and checks that the health check passes, that a
|
`deploy/example-app`, runs it with a volume for the state files, and checks
|
||||||
request reaches the app through `smallwebwaf`, and that `sv stop` and
|
that the health check passes, that a request reaches the app through
|
||||||
`docker stop` stop it in order; then removes the container and both images. It
|
`smallwebwaf`, that a second request in a minute bans the client, that a probe
|
||||||
needs network access, for nixpkgs' binary cache, and `script/check` does not
|
for `/.env` bans another client, whose next request makes the ban permanent,
|
||||||
run it; `make example-app` does.
|
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
|
## TODO
|
||||||
|
|
||||||
- The rest of milestone 3, after the static lists, and the rest of the design,
|
- The rest of the design, in the order of the build order in
|
||||||
in the order of the build order in [`SPEC.md`](SPEC.md).
|
[`SPEC.md`](SPEC.md).
|
||||||
|
|
||||||
## Documents
|
## Documents
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from a directory of hand-editable text files.
|
|||||||
- Defence against traffic floods that saturate the host's network link. That
|
- Defence against traffic floods that saturate the host's network link. That
|
||||||
needs help upstream of the host.
|
needs help upstream of the host.
|
||||||
- A web UI or a configuration file. Settings are environment variables. Apart
|
- A web UI or a configuration file. Settings are environment variables. Apart
|
||||||
from settings given as files (the `_FILE` form of any setting, such as
|
from settings given as files (the `_FILE` form of a setting, such as
|
||||||
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
|
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
|
||||||
files and the lookup database, the only files read are the rule files, which
|
files and the lookup database, the only files read are the rule files, which
|
||||||
hold one regex per line and nothing more elaborate.
|
hold one regex per line and nothing more elaborate.
|
||||||
@@ -293,11 +293,13 @@ it.
|
|||||||
needs: an alert destination, an account key, a token.
|
needs: an alert destination, an account key, a token.
|
||||||
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
|
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
|
||||||
container, and so its environment variables, with the app it protects.
|
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.
|
- 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
|
- 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
|
setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for
|
||||||
secrets and long lists.
|
secrets and long lists. `SWWAF_LOG_REMOTE_TLS_CA_FILE`, whose value names a
|
||||||
|
file already, has no `_FILE` form.
|
||||||
- Settings, including those given as files, are read once at start; changing one
|
- Settings, including those given as files, are read once at start; changing one
|
||||||
means restarting the container. The files `smallwebwaf` watches while it runs
|
means restarting the container. The files `smallwebwaf` watches while it runs
|
||||||
are its state files, its rule files and the lookup database.
|
are its state files, its rule files and the lookup database.
|
||||||
@@ -413,7 +415,9 @@ The settings, by group:
|
|||||||
headers, its body.
|
headers, its body.
|
||||||
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
|
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
|
||||||
request line and headers a client may send. Over it, `smallwebwaf` answers
|
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
|
- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open
|
||||||
connection may wait for its next request before `smallwebwaf` closes it.
|
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
|
It is longer than the 90 seconds after which traefik, by default, closes a
|
||||||
@@ -945,9 +949,9 @@ and the running `smallwebwaf` takes the edit in.
|
|||||||
- what was broken: the rule ids and target that matched, or the limit, its
|
- 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
|
window, the count reached and the client's limit percentage with what set
|
||||||
it; and any reputation sources that listed the client;
|
it; and any reputation sources that listed the client;
|
||||||
- the requests that caused the ban, up to the last ten: time, method, host,
|
- the request that caused the ban, the one that broke the limit or carried
|
||||||
path with its query string, status and user agent, each text cut to 256
|
the clear sign of attack: time, method, host, path with its query string,
|
||||||
bytes;
|
status and user agent, each text cut to 256 bytes;
|
||||||
- how many requests counted toward the ban, and the time span over which
|
- how many requests counted toward the ban, and the time span over which
|
||||||
they came;
|
they came;
|
||||||
- the netblock's total requests since it was first seen, and the requests
|
- the netblock's total requests since it was first seen, and the requests
|
||||||
@@ -958,13 +962,13 @@ and the running `smallwebwaf` takes the edit in.
|
|||||||
the table is full, so on a public service the file grows to the default
|
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
|
`SWWAF_MAX_TRACKED_CLIENTS` of 20,000, about 20 MiB. Written every 15
|
||||||
minutes, that is under 2 GiB of disk writes a day.
|
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
|
- `bans.json` takes about 1.2 KiB per ban and at most about 2.5 KiB, since
|
||||||
texts in the notes are cut short. At the default `SWWAF_MAX_BANS` of 5,000
|
the notes hold one request and their texts are cut short. At the default
|
||||||
it is about 10 MiB, and never more than about 40 MiB, plus whatever bans
|
`SWWAF_MAX_BANS` of 5,000 it is about 6 MiB, and never more than about 12
|
||||||
an admin made. It is written when a ban is made, lifted or made permanent,
|
MiB, plus whatever bans an admin made. It is written when a ban is made,
|
||||||
at most once every 10 seconds, and otherwise with the 15-minute write, so
|
lifted or made permanent, at most once every 10 seconds, and otherwise
|
||||||
its writes follow the bans made: with a full file, a hundred new bans a
|
with the 15-minute write, so its writes follow the bans made: with a full
|
||||||
day come to about 1 GiB of disk writes.
|
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.
|
- `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.
|
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.
|
- `reputation.json` and `alerts.json` are usually a few MiB or less.
|
||||||
|
|||||||
@@ -2,4 +2,20 @@ module sneak.berlin/go/smallwebwaf
|
|||||||
|
|
||||||
go 1.26.0
|
go 1.26.0
|
||||||
|
|
||||||
require github.com/hashicorp/golang-lru/v2 v2.0.7
|
require (
|
||||||
|
github.com/fsnotify/fsnotify v1.10.1
|
||||||
|
github.com/hashicorp/golang-lru/v2 v2.0.7
|
||||||
|
github.com/prometheus/client_golang v1.24.1
|
||||||
|
)
|
||||||
|
|
||||||
|
require (
|
||||||
|
github.com/beorn7/perks v1.0.1 // indirect
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
|
github.com/kylelemons/godebug v1.1.0 // indirect
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
|
github.com/prometheus/common v0.70.1 // indirect
|
||||||
|
github.com/prometheus/procfs v0.21.1 // indirect
|
||||||
|
golang.org/x/sys v0.47.0 // indirect
|
||||||
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,2 +1,40 @@
|
|||||||
|
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||||
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
|
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||||
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
|
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||||
|
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||||
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
|
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
||||||
|
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
|
||||||
|
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
|
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||||
|
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||||
|
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
||||||
|
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
||||||
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
|
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||||
|
github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY=
|
||||||
|
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
||||||
|
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
||||||
|
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||||
|
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||||
|
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||||
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
|
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
|
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
|
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||||
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
|
|||||||
@@ -0,0 +1,186 @@
|
|||||||
|
package bans_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBanWithoutACauseIsAnAdmins(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{{Netblock: netblock, Start: midnight()}})
|
||||||
|
|
||||||
|
if got := ledger.Bans(netblock)[0].Cause; got != bans.CauseAdmin {
|
||||||
|
t.Errorf("the ban's cause is %q, want admin", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAdminsBansAreNeverDroppedAndDoNotCountTowardMaxBans(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
rules := defaultRules()
|
||||||
|
rules.MaxBans = 1
|
||||||
|
ledger := bans.New(rules)
|
||||||
|
adminsOnly := netip.MustParsePrefix("198.51.100.0/24")
|
||||||
|
both := netip.MustParsePrefix("203.0.113.1/32")
|
||||||
|
second := netip.MustParsePrefix("203.0.113.2/32")
|
||||||
|
third := netip.MustParsePrefix("203.0.113.3/32")
|
||||||
|
|
||||||
|
// Seen longest ago, a netblock with two of an admin's bans alone, and
|
||||||
|
// then one with an admin's ban before a ban smallwebwaf made: the one
|
||||||
|
// ban counted toward MaxBans.
|
||||||
|
ledger.Load([]bans.Ban{
|
||||||
|
{Netblock: adminsOnly, Start: midnight().Add(-3 * time.Hour), Cause: bans.CauseAdmin},
|
||||||
|
{Netblock: adminsOnly, Start: midnight().Add(-2 * time.Hour), Cause: bans.CauseAdmin},
|
||||||
|
{Netblock: both, Start: midnight().Add(-time.Hour), Cause: bans.CauseAdmin},
|
||||||
|
{
|
||||||
|
Netblock: both,
|
||||||
|
Start: midnight(),
|
||||||
|
Expires: midnight().Add(time.Hour),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 2})
|
||||||
|
|
||||||
|
// A new ban drops the ban smallwebwaf made, and only that one.
|
||||||
|
ledger.BanForLimit(second, midnight(), bans.Notes{})
|
||||||
|
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 1})
|
||||||
|
|
||||||
|
if ledger.Bans(both)[0].Cause != bans.CauseAdmin {
|
||||||
|
t.Errorf("%s kept %+v, want the admin's ban", both, ledger.Bans(both))
|
||||||
|
}
|
||||||
|
|
||||||
|
// And the next drops that one.
|
||||||
|
ledger.BanForLimit(third, midnight(), bans.Notes{})
|
||||||
|
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 0, third: 1})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
|
||||||
|
limit := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
||||||
|
bans.Notes{Limit: 1000, Window: "minute"})
|
||||||
|
attack := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
|
||||||
|
bans.Notes{RuleID: "git-dir", Target: "path"})
|
||||||
|
|
||||||
|
for _, tc := range []struct{ got, want string }{
|
||||||
|
{limit.Reason, "requests per minute over the limit of 1000"},
|
||||||
|
{attack.Reason, "matched the rule git-dir"},
|
||||||
|
} {
|
||||||
|
if tc.got != tc.want {
|
||||||
|
t.Errorf("the reason is %q, want %q", tc.got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// An hour's ban lifted ten minutes after it started.
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
lifted := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: midnight(),
|
||||||
|
Expires: midnight().Add(time.Hour),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
Lifted: midnight().Add(10 * time.Minute),
|
||||||
|
}
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{lifted})
|
||||||
|
|
||||||
|
// While it would still last, it refuses nothing, and a limit broken
|
||||||
|
// bans for an hour, as a first broken limit does; the lifted ban is
|
||||||
|
// kept, and counted among the earlier bans.
|
||||||
|
now := midnight().Add(30 * time.Minute)
|
||||||
|
|
||||||
|
_, banned := ledger.Check(netblock.Addr(), now)
|
||||||
|
if banned {
|
||||||
|
t.Error("the lifted ban refuses")
|
||||||
|
}
|
||||||
|
|
||||||
|
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
|
if ban.Expires.Sub(ban.Start) != time.Hour ||
|
||||||
|
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
||||||
|
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
|
||||||
|
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
|
||||||
|
}
|
||||||
|
|
||||||
|
held := ledger.Bans(netblock)
|
||||||
|
if len(held) != 2 || held[0] != lifted {
|
||||||
|
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// A permanent ban for a clear sign of attack, lifted.
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: midnight(),
|
||||||
|
Cause: bans.CauseAttack,
|
||||||
|
Lifted: midnight().Add(time.Hour),
|
||||||
|
}})
|
||||||
|
|
||||||
|
now := midnight().Add(2 * time.Hour)
|
||||||
|
|
||||||
|
_, banned := ledger.Find(netblock.Addr(), now)
|
||||||
|
if banned {
|
||||||
|
t.Error("the lifted ban refuses")
|
||||||
|
}
|
||||||
|
|
||||||
|
active, permanent := ledger.Count(now)
|
||||||
|
if active != 0 || permanent != 0 {
|
||||||
|
t.Errorf("%d bans are active and %d permanent, want none", active, permanent)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next clear sign of attack bans for seven days, as a first does.
|
||||||
|
ban := ledger.BanForAttack(netblock, now, bans.Notes{})
|
||||||
|
if ban.Expires.Sub(ban.Start) != 7*day {
|
||||||
|
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
made := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
||||||
|
bans.Notes{})
|
||||||
|
atStart := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
||||||
|
Start: midnight(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// The bans read at the start were made before it.
|
||||||
|
ledger.Load([]bans.Ban{made, atStart})
|
||||||
|
|
||||||
|
if got := ledger.Made(bans.CauseAdmin); got != 0 {
|
||||||
|
t.Fatalf("%d bans made by an admin after the start's, want none", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The admin keeps the ban smallwebwaf made, keeps the one read at the
|
||||||
|
// start, and adds one without a cause: that one alone is made.
|
||||||
|
kept := made
|
||||||
|
kept.Cause = bans.CauseAdmin
|
||||||
|
added := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.3/32"),
|
||||||
|
Start: midnight(),
|
||||||
|
}
|
||||||
|
ledger.LoadEdit([]bans.Ban{kept, atStart, added})
|
||||||
|
|
||||||
|
if ledger.Made(bans.CauseAdmin) != 1 || ledger.Made(bans.CauseLimit) != 1 {
|
||||||
|
t.Errorf("%d bans made by an admin and %d for a limit, want 1 of each",
|
||||||
|
ledger.Made(bans.CauseAdmin), ledger.Made(bans.CauseLimit))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,631 @@
|
|||||||
|
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
||||||
|
// netblocks of clients that break a rate limit or show a clear sign of
|
||||||
|
// attack, and those an admin makes, 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 (
|
||||||
|
"fmt"
|
||||||
|
"math"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The causes of bans.
|
||||||
|
const (
|
||||||
|
// CauseLimit is a ban smallwebwaf made for a broken limit.
|
||||||
|
CauseLimit = "limit"
|
||||||
|
// CauseAttack is a ban smallwebwaf made for a clear sign of attack.
|
||||||
|
CauseAttack = "attack"
|
||||||
|
// CauseAdmin is a ban an admin made, or one smallwebwaf made that an
|
||||||
|
// admin keeps. It is never dropped.
|
||||||
|
CauseAdmin = "admin"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 lasts, and how many bans are held.
|
||||||
|
type Rules struct {
|
||||||
|
// LimitBanDuration is how long a first ban for a broken limit lasts.
|
||||||
|
LimitBanDuration time.Duration
|
||||||
|
// LimitBanRepeatWindow is how soon after the end of the netblock's
|
||||||
|
// ban that ended last, other than one for a clear sign of attack, 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 for a broken limit; one that would
|
||||||
|
// be longer is permanent instead.
|
||||||
|
MaxBanDuration time.Duration
|
||||||
|
// AttackBanDuration is how long a first ban for a clear sign of attack
|
||||||
|
// lasts.
|
||||||
|
AttackBanDuration time.Duration
|
||||||
|
// MaxBans is the most bans held whose cause is not CauseAdmin, at
|
||||||
|
// least one. Past it, the earliest such ban of the netblock that has
|
||||||
|
// gone longest without a request is dropped. Bans whose cause is
|
||||||
|
// CauseAdmin are held besides, and never dropped.
|
||||||
|
MaxBans int
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ban is a ban on a netblock.
|
||||||
|
type Ban struct {
|
||||||
|
Netblock netip.Prefix
|
||||||
|
Start time.Time
|
||||||
|
// Expires is when the ban ends, zero for a permanent ban.
|
||||||
|
Expires time.Time
|
||||||
|
// Cause is CauseLimit, CauseAttack or CauseAdmin.
|
||||||
|
Cause string
|
||||||
|
// Reason is a short text: for a ban smallwebwaf made, the limit broken
|
||||||
|
// or the rule that matched; for an admin's, what the admin wrote.
|
||||||
|
Reason string
|
||||||
|
// Lifted is when an admin lifted the ban, zero while no admin has. A
|
||||||
|
// lifted ban refuses nothing, and does not make the netblock's next
|
||||||
|
// ban longer.
|
||||||
|
Lifted 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: it has not
|
||||||
|
// been lifted, and has not run out.
|
||||||
|
func (b Ban) ActiveAt(now time.Time) bool {
|
||||||
|
return b.Lifted.IsZero() && (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, for a ban for a broken limit, 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,omitempty"`
|
||||||
|
Window string `json:"window,omitempty"`
|
||||||
|
Count float64 `json:"count,omitempty"`
|
||||||
|
// RuleID and Target are, for a ban for a clear sign of attack, the id
|
||||||
|
// of the rule file rule that matched, and its target.
|
||||||
|
RuleID string `json:"rule_id,omitempty"`
|
||||||
|
Target string `json:"target,omitempty"`
|
||||||
|
// Request is the request that broke the limit, or that was the clear
|
||||||
|
// sign of attack.
|
||||||
|
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, by
|
||||||
|
// cause.
|
||||||
|
EarlierBans EarlierBans `json:"earlier_bans"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// EarlierBans counts a netblock's bans before a ban, by cause.
|
||||||
|
type EarlierBans struct {
|
||||||
|
Limit int `json:"limit"`
|
||||||
|
Attack int `json:"attack"`
|
||||||
|
Admin int `json:"admin"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 and
|
||||||
|
// Find make each netblock they find the most recently seen.
|
||||||
|
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
|
||||||
|
// held is how many bans netblocks holds whose cause is not CauseAdmin,
|
||||||
|
// at most rules.MaxBans.
|
||||||
|
held int
|
||||||
|
// made is how many bans have been made since the start, by cause: by
|
||||||
|
// the ledger, and by an admin in an edit of bans.json.
|
||||||
|
made map[string]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 {
|
||||||
|
// The ledger drops bans itself, and never those whose cause is
|
||||||
|
// CauseAdmin, however many there are, so the LRU has no limit of its
|
||||||
|
// own: it keeps the netblocks in the order they were last seen.
|
||||||
|
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](math.MaxInt, 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,
|
||||||
|
made: map[string]int{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Changed receives a value after a ban is made or made permanent, so that
|
||||||
|
// bans.json can be written. Several changes before it is read leave one
|
||||||
|
// value.
|
||||||
|
func (l *Ledger) Changed() <-chan struct{} {
|
||||||
|
return l.changed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check is called for a 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. A ban for a clear sign of
|
||||||
|
// attack is made permanent by the request: the netblock is malicious.
|
||||||
|
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
ban := l.active(client, now)
|
||||||
|
if ban == nil {
|
||||||
|
return Ban{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
ban.Notes.Requests++
|
||||||
|
ban.Notes.Refused++
|
||||||
|
|
||||||
|
if ban.Cause == CauseAttack && !ban.Permanent() {
|
||||||
|
ban.Expires = time.Time{}
|
||||||
|
|
||||||
|
l.markChanged()
|
||||||
|
}
|
||||||
|
|
||||||
|
return *ban, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find is Check without counting the request among those the ban
|
||||||
|
// refused: in observe mode a ban refuses nothing.
|
||||||
|
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
ban := l.active(client, now)
|
||||||
|
if ban == nil {
|
||||||
|
return Ban{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return *ban, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// activeBan returns the ban in bans, a netblock's bans oldest first, that
|
||||||
|
// is active at now, or nil when none is. If several are, it returns the
|
||||||
|
// one that started last. Every ban is looked at, since a ban an admin adds
|
||||||
|
// to bans.json can start before the netblock's others and outlast them.
|
||||||
|
func activeBan(bans []Ban, now time.Time) *Ban {
|
||||||
|
for i := len(bans) - 1; i >= 0; i-- {
|
||||||
|
if bans[i].ActiveAt(now) {
|
||||||
|
return &bans[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// BanForLimit bans netblock at now for a broken limit, with notes, and
|
||||||
|
// returns the ban. A first ban lasts LimitBanDuration. A ban made within
|
||||||
|
// LimitBanRepeatWindow after the netblock's ban that ended last, other
|
||||||
|
// than one for a clear sign of attack or a lifted one, 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, and gives the ban the reason "requests per
|
||||||
|
// <Window> over the limit of <Limit>", from the notes.
|
||||||
|
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
|
||||||
|
reason := fmt.Sprintf("requests per %s over the limit of %d",
|
||||||
|
notes.Window, notes.Limit)
|
||||||
|
|
||||||
|
return l.ban(netblock, now, CauseLimit, reason, notes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BanForAttack bans netblock at now for a clear sign of attack, with
|
||||||
|
// notes, and returns the ban, as BanForLimit does. A first ban lasts
|
||||||
|
// AttackBanDuration; once the netblock has had one that was not lifted,
|
||||||
|
// the next is permanent. Its reason is "matched the rule <RuleID>".
|
||||||
|
func (l *Ledger) BanForAttack(netblock netip.Prefix, now time.Time, notes Notes) Ban {
|
||||||
|
return l.ban(netblock, now, CauseAttack, "matched the rule "+notes.RuleID, notes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Made returns how many bans for cause have been made since the start:
|
||||||
|
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
|
||||||
|
// admin in an edit of bans.json, as LoadEdit counts them. The bans read
|
||||||
|
// from bans.json at the start are not among them.
|
||||||
|
func (l *Ledger) Made(cause string) int {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
return l.made[cause]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Count returns how many of the bans held are active at now, and how many
|
||||||
|
// of those are permanent. A lifted ban is neither.
|
||||||
|
func (l *Ledger) Count(now time.Time) (int, int) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
active, permanent := 0, 0
|
||||||
|
|
||||||
|
for _, bans := range l.netblocks.Values() {
|
||||||
|
for _, ban := range *bans {
|
||||||
|
if !ban.ActiveAt(now) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
active++
|
||||||
|
|
||||||
|
if ban.Permanent() {
|
||||||
|
permanent++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return active, permanent
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 at the start into the ledger, in
|
||||||
|
// place of the bans it holds, in the order they started, so that a
|
||||||
|
// netblock whose last ban started latest counts as the most recently
|
||||||
|
// seen. A ban without a cause is an admin's, and gets CauseAdmin. 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 whose cause is not CauseAdmin are dropped, as
|
||||||
|
// when they are made.
|
||||||
|
func (l *Ledger) Load(bans []Ban) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
l.load(bans)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadEdit is Load for an admin's edit of bans.json, taken in while
|
||||||
|
// smallwebwaf runs. Each ban in it whose cause is CauseAdmin, and which
|
||||||
|
// the ledger did not hold, with the same netblock and start, is one the
|
||||||
|
// admin made, and is counted among the bans made.
|
||||||
|
func (l *Ledger) LoadEdit(bans []Ban) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
l.made[CauseAdmin] += l.load(bans)
|
||||||
|
}
|
||||||
|
|
||||||
|
// load does what Load describes, and returns how many of bans are bans
|
||||||
|
// whose cause is CauseAdmin that the ledger did not hold before.
|
||||||
|
func (l *Ledger) load(bans []Ban) int {
|
||||||
|
bans = slices.Clone(bans)
|
||||||
|
added := 0
|
||||||
|
|
||||||
|
for i := range bans {
|
||||||
|
ban := &bans[i]
|
||||||
|
ban.Netblock = ban.Netblock.Masked()
|
||||||
|
ban.Notes.Request = ban.Notes.Request.cut()
|
||||||
|
|
||||||
|
if ban.Cause == "" {
|
||||||
|
ban.Cause = CauseAdmin
|
||||||
|
}
|
||||||
|
|
||||||
|
if ban.Cause == CauseAdmin && !l.holds(ban.Netblock, ban.Start) {
|
||||||
|
added++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
slices.SortStableFunc(bans, func(a, b Ban) int {
|
||||||
|
return a.Start.Compare(b.Start)
|
||||||
|
})
|
||||||
|
|
||||||
|
l.netblocks.Purge()
|
||||||
|
l.held = 0
|
||||||
|
l.v4Lengths, l.v6Lengths = nil, nil
|
||||||
|
|
||||||
|
for _, ban := range bans {
|
||||||
|
l.add(ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
return added
|
||||||
|
}
|
||||||
|
|
||||||
|
// holds reports whether the ledger holds a ban on netblock that started
|
||||||
|
// at start.
|
||||||
|
func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
|
||||||
|
bans, found := l.netblocks.Peek(netblock)
|
||||||
|
|
||||||
|
return found && slices.ContainsFunc(*bans, func(ban Ban) bool {
|
||||||
|
return ban.Start.Equal(start)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// ban bans netblock at now for cause, with reason and notes, as
|
||||||
|
// BanForLimit and BanForAttack describe, and returns the ban.
|
||||||
|
func (l *Ledger) ban(
|
||||||
|
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes,
|
||||||
|
) Ban {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
// held are the netblock's bans, none of them active.
|
||||||
|
var held []Ban
|
||||||
|
|
||||||
|
bans, found := l.netblocks.Get(netblock)
|
||||||
|
if found {
|
||||||
|
active := activeBan(*bans, now)
|
||||||
|
if active != nil {
|
||||||
|
return *active
|
||||||
|
}
|
||||||
|
|
||||||
|
held = *bans
|
||||||
|
notes.EarlierBans = earlierBans(held)
|
||||||
|
}
|
||||||
|
|
||||||
|
notes.Request = notes.Request.cut()
|
||||||
|
ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes}
|
||||||
|
|
||||||
|
if cause == CauseAttack {
|
||||||
|
ban.Expires = l.attackExpiry(held, now)
|
||||||
|
} else {
|
||||||
|
ban.Expires = l.limitExpiry(held, now)
|
||||||
|
}
|
||||||
|
|
||||||
|
l.add(ban)
|
||||||
|
l.made[cause]++
|
||||||
|
l.markChanged()
|
||||||
|
|
||||||
|
return ban
|
||||||
|
}
|
||||||
|
|
||||||
|
// earlierBans returns how many bans a netblock with the bans held, oldest
|
||||||
|
// first, has had, by cause: the first ban held counts the bans the
|
||||||
|
// netblock had before that one, since dropped to make room, and each ban
|
||||||
|
// held adds one.
|
||||||
|
func earlierBans(held []Ban) EarlierBans {
|
||||||
|
earlier := held[0].Notes.EarlierBans
|
||||||
|
|
||||||
|
for _, ban := range held {
|
||||||
|
switch ban.Cause {
|
||||||
|
case CauseLimit:
|
||||||
|
earlier.Limit++
|
||||||
|
case CauseAttack:
|
||||||
|
earlier.Attack++
|
||||||
|
case CauseAdmin:
|
||||||
|
earlier.Admin++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return earlier
|
||||||
|
}
|
||||||
|
|
||||||
|
// markChanged has Changed receive a value, unless one is waiting already.
|
||||||
|
func (l *Ledger) markChanged() {
|
||||||
|
select {
|
||||||
|
case l.changed <- struct{}{}:
|
||||||
|
default: // a value is waiting already
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// active returns the ban active at now on a netblock client is in, or
|
||||||
|
// nil.
|
||||||
|
func (l *Ledger) active(client netip.Addr, now time.Time) *Ban {
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
ban := activeBan(*bans, now)
|
||||||
|
if ban != nil {
|
||||||
|
return ban
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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,
|
||||||
|
// unless ban's cause is CauseAdmin, which does not count toward MaxBans.
|
||||||
|
func (l *Ledger) add(ban Ban) {
|
||||||
|
counted := ban.Cause != CauseAdmin
|
||||||
|
if counted && 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)
|
||||||
|
|
||||||
|
if counted {
|
||||||
|
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())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// limitExpiry returns when a ban for a broken limit made at now ends, or
|
||||||
|
// zero when it is permanent. held are the netblock's bans, none of them
|
||||||
|
// active, of which the one that ended last, other than a ban for a clear
|
||||||
|
// sign of attack or a lifted one, can make the new ban longer. A ban an
|
||||||
|
// admin adds to bans.json can start after another and end before it, so
|
||||||
|
// that one is looked for among them all.
|
||||||
|
func (l *Ledger) limitExpiry(held []Ban, now time.Time) time.Time {
|
||||||
|
length := l.rules.LimitBanDuration
|
||||||
|
|
||||||
|
var last *Ban
|
||||||
|
|
||||||
|
for i, ban := range held {
|
||||||
|
if ban.Cause != CauseAttack && ban.Lifted.IsZero() &&
|
||||||
|
(last == nil || ban.Expires.After(last.Expires)) {
|
||||||
|
last = &held[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
// attackExpiry returns when a ban for a clear sign of attack made at now
|
||||||
|
// ends. held are the netblock's bans, none of them active: if one of them
|
||||||
|
// is for a clear sign of attack too, and was not lifted, the new ban is
|
||||||
|
// permanent, and its end zero; otherwise it ends AttackBanDuration later.
|
||||||
|
func (l *Ledger) attackExpiry(held []Ban, now time.Time) time.Time {
|
||||||
|
for _, ban := range held {
|
||||||
|
if ban.Cause == CauseAttack && ban.Lifted.IsZero() {
|
||||||
|
return time.Time{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return now.Add(l.rules.AttackBanDuration)
|
||||||
|
}
|
||||||
|
|
||||||
|
// dropOne drops the earliest ban whose cause is not CauseAdmin of the
|
||||||
|
// netblock that has gone longest without a request, of those that hold
|
||||||
|
// such a ban, and the netblock with it if that was its only ban. It is
|
||||||
|
// called with at least one such ban held. It looks at each netblock once
|
||||||
|
// at most, and drops nothing when none holds such a ban.
|
||||||
|
func (l *Ledger) dropOne() {
|
||||||
|
for range l.netblocks.Len() {
|
||||||
|
netblock, bans, _ := l.netblocks.GetOldest()
|
||||||
|
|
||||||
|
i := slices.IndexFunc(*bans, func(ban Ban) bool {
|
||||||
|
return ban.Cause != CauseAdmin
|
||||||
|
})
|
||||||
|
if i < 0 {
|
||||||
|
// Its bans are all an admin's, and never dropped. Get makes
|
||||||
|
// it the most recently seen, so that the next netblock is
|
||||||
|
// looked at; when it was seen matters only for dropping a
|
||||||
|
// ban, and a ban added to it makes it the most recently seen
|
||||||
|
// anyway.
|
||||||
|
l.netblocks.Get(netblock)
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(*bans) == 1 {
|
||||||
|
l.netblocks.Remove(netblock)
|
||||||
|
} else {
|
||||||
|
*bans = slices.Delete(*bans, i, i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
l.held--
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)])
|
||||||
|
}
|
||||||
@@ -0,0 +1,387 @@
|
|||||||
|
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 != (bans.EarlierBans{Limit: i}) {
|
||||||
|
t.Fatalf("ban %d lasts %s with earlier bans %+v, want %d hours and %d for a limit",
|
||||||
|
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 != (bans.EarlierBans{Limit: 1}) {
|
||||||
|
t.Errorf("second ban lasts %s with earlier bans %+v, want %s and 1 for a limit",
|
||||||
|
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 TestFindCountsNothing(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})
|
||||||
|
|
||||||
|
got, banned := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
||||||
|
if !banned || got != ban {
|
||||||
|
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, banned = ledger.Find(netblock.Addr(), ban.Expires)
|
||||||
|
if banned {
|
||||||
|
t.Error("the ban did not end")
|
||||||
|
}
|
||||||
|
|
||||||
|
if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes {
|
||||||
|
t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != (bans.EarlierBans{Limit: 1}) {
|
||||||
|
t.Errorf("the ledger holds %+v, want only the second ban, "+
|
||||||
|
"with 1 earlier ban for a limit", held)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
notes := bans.Notes{RuleID: "env-file", Target: "path"}
|
||||||
|
|
||||||
|
ban := ledger.BanForAttack(netblock, midnight(), notes)
|
||||||
|
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
|
||||||
|
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
|
||||||
|
ledger.Made(bans.CauseLimit) != 0 {
|
||||||
|
t.Fatalf("the ban is %+v, with %d made for an attack and %d for a limit, "+
|
||||||
|
"want one for an attack, of seven days", ban,
|
||||||
|
ledger.Made(bans.CauseAttack), ledger.Made(bans.CauseLimit))
|
||||||
|
}
|
||||||
|
|
||||||
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
|
// In observe mode the ban refuses nothing, and stays as it is.
|
||||||
|
got, _ := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
|
||||||
|
if got.Permanent() {
|
||||||
|
t.Fatal("a request found under the ban made it permanent")
|
||||||
|
}
|
||||||
|
|
||||||
|
// A request it refuses makes it permanent, and bans.json due.
|
||||||
|
got, _ = ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
|
||||||
|
if !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
|
||||||
|
t.Fatalf("after a request during the ban, it is %+v, want it permanent", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantChanged(t, ledger, true)
|
||||||
|
|
||||||
|
_, banned := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
|
||||||
|
if !banned {
|
||||||
|
t.Error("the permanent ban ended")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
// A ban for a broken limit before does not count.
|
||||||
|
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
||||||
|
second := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
|
||||||
|
|
||||||
|
if second.Expires.Sub(second.Start) != 7*day {
|
||||||
|
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
|
||||||
|
second.Expires.Sub(second.Start))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Once that has run out without a request, the netblock is served, and
|
||||||
|
// its next clear sign of attack bans it for good.
|
||||||
|
_, banned := ledger.Check(netblock.Addr(), second.Expires)
|
||||||
|
if banned {
|
||||||
|
t.Fatal("the ban did not end")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Its notes show the earlier ban for an attack that makes it permanent,
|
||||||
|
// beside the one for a limit.
|
||||||
|
third := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
|
||||||
|
if !third.Permanent() ||
|
||||||
|
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
|
||||||
|
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
|
||||||
|
"with 1 earlier ban for a limit and 1 for an attack", third)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
|
// Three times the seven days would be permanent; a limit broken as the
|
||||||
|
// ban for an attack ends bans for an hour, as a first broken limit does.
|
||||||
|
attack := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
|
||||||
|
limit := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
|
||||||
|
|
||||||
|
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
|
||||||
|
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
|
||||||
|
}
|
||||||
|
|
||||||
|
// And a request during the ban for a limit leaves it as it is.
|
||||||
|
got, _ := ledger.Check(netblock.Addr(), limit.Start)
|
||||||
|
if got.Permanent() {
|
||||||
|
t.Error("a request during a ban for a limit made it permanent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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,
|
||||||
|
AttackBanDuration: 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,314 @@
|
|||||||
|
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 != (bans.EarlierBans{Limit: 1}) {
|
||||||
|
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
|
||||||
|
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 TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// As when an admin adds a permanent ban to bans.json with a start
|
||||||
|
// before that of the netblock's ban that has ended.
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
||||||
|
permanent := bans.Ban{Netblock: netblock, Start: midnight().Add(-time.Hour)}
|
||||||
|
ended := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: midnight(),
|
||||||
|
Expires: midnight().Add(time.Hour),
|
||||||
|
}
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{permanent, ended})
|
||||||
|
|
||||||
|
now := midnight().Add(2 * time.Hour)
|
||||||
|
client := netip.MustParseAddr("203.0.113.9")
|
||||||
|
|
||||||
|
ban, banned := ledger.Find(client, now)
|
||||||
|
if !banned || !ban.Permanent() {
|
||||||
|
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
|
||||||
|
}
|
||||||
|
|
||||||
|
ban, banned = ledger.Check(client, now)
|
||||||
|
if !banned || !ban.Permanent() {
|
||||||
|
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
|
||||||
|
banned, ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A limit broken now makes no shorter ban over the permanent one.
|
||||||
|
ban = ledger.BanForLimit(netblock, now, bans.Notes{})
|
||||||
|
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
|
||||||
|
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
|
||||||
|
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// A 9-hour ban smallwebwaf made, the third in a row, and an admin's
|
||||||
|
// 1-hour ban added to bans.json over it, with no cause and no notes.
|
||||||
|
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
nineHours := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: midnight(),
|
||||||
|
Expires: midnight().Add(9 * time.Hour),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 2}},
|
||||||
|
}
|
||||||
|
admins := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: midnight().Add(time.Hour),
|
||||||
|
Expires: midnight().Add(2 * time.Hour),
|
||||||
|
}
|
||||||
|
|
||||||
|
ledger := bans.New(defaultRules())
|
||||||
|
ledger.Load([]bans.Ban{nineHours, admins})
|
||||||
|
|
||||||
|
// Once both have ended, a limit broken within the repeat window bans
|
||||||
|
// for three times the 9 hours, and the notes count the two bans
|
||||||
|
// before the 9-hour one and it, for a limit, and the admin's.
|
||||||
|
ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
|
||||||
|
if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
|
||||||
|
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
|
||||||
|
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
|
||||||
|
"want 27h, 3 for a limit and 1 an admin's",
|
||||||
|
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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(),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
}
|
||||||
|
earlier := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
||||||
|
Start: midnight().Add(-time.Hour),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
}
|
||||||
|
|
||||||
|
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 TestLoadReplacesTheBansHeld(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Room for three bans, so that the second load, were it added to the
|
||||||
|
// two bans held, would drop none of them to make room.
|
||||||
|
rules := defaultRules()
|
||||||
|
rules.MaxBans = 3
|
||||||
|
ledger := bans.New(rules)
|
||||||
|
kept := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix("2001:db8::/64"),
|
||||||
|
Start: midnight(),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
}
|
||||||
|
ledger.Load([]bans.Ban{
|
||||||
|
{
|
||||||
|
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
|
||||||
|
Start: midnight(),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
},
|
||||||
|
kept,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Loaded again without the first ban, as when an admin's edit of
|
||||||
|
// bans.json is taken in, that ban is lifted.
|
||||||
|
ledger.Load([]bans.Ban{kept})
|
||||||
|
|
||||||
|
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
|
||||||
|
if banned {
|
||||||
|
t.Error("a ban left out of the second load still refuses")
|
||||||
|
}
|
||||||
|
|
||||||
|
// The ledger holds one ban, so it makes two more without dropping any.
|
||||||
|
first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
|
||||||
|
bans.Notes{})
|
||||||
|
second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
|
||||||
|
bans.Notes{})
|
||||||
|
|
||||||
|
want := []bans.Ban{first, second, kept}
|
||||||
|
if got := ledger.Snapshot(); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("the ledger holds %+v, want %+v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+592
-13
@@ -1,20 +1,28 @@
|
|||||||
// Package config reads smallwebwaf's settings. Every setting is an
|
// Package config reads smallwebwaf's settings. Every setting is an
|
||||||
// environment variable whose name starts with SWWAF_, every setting has a
|
// environment variable whose name starts with SWWAF_, or a file such a
|
||||||
// default, and this package is the one place they are read.
|
// variable names, every setting has a default, and this package is the
|
||||||
|
// one place they are read.
|
||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/x509"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"math"
|
"math"
|
||||||
"net"
|
"net"
|
||||||
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"slices"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
|
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
|
||||||
@@ -24,12 +32,28 @@ type Config struct {
|
|||||||
ListenAddr string
|
ListenAddr string
|
||||||
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
|
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
|
||||||
UpstreamURL *url.URL
|
UpstreamURL *url.URL
|
||||||
|
// InstanceName is the name each request log line gives as instance
|
||||||
|
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets
|
||||||
|
// to the first 12 characters of the container's id.
|
||||||
|
InstanceName string
|
||||||
|
// Observe is true in observe mode, when SWWAF_MODE is observe rather
|
||||||
|
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country
|
||||||
|
// lists, a rate limit or a rule would refuse is passed to the app
|
||||||
|
// instead, and no ban is made.
|
||||||
|
Observe bool
|
||||||
// TrustedProxies are the netblocks whose X-Forwarded-For is
|
// TrustedProxies are the netblocks whose X-Forwarded-For is
|
||||||
// believed (SWWAF_TRUSTED_PROXIES).
|
// believed (SWWAF_TRUSTED_PROXIES).
|
||||||
TrustedProxies []netip.Prefix
|
TrustedProxies []netip.Prefix
|
||||||
// ClientRequestTimeout bounds reading the whole request from the
|
// ClientRequestTimeout bounds reading the whole request from the
|
||||||
// client (SWWAF_CLIENT_REQUEST_TIMEOUT).
|
// client (SWWAF_CLIENT_REQUEST_TIMEOUT).
|
||||||
ClientRequestTimeout time.Duration
|
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
|
// ClientResponseTimeout bounds writing the whole response to the
|
||||||
// client (SWWAF_CLIENT_RESPONSE_TIMEOUT).
|
// client (SWWAF_CLIENT_RESPONSE_TIMEOUT).
|
||||||
ClientResponseTimeout time.Duration
|
ClientResponseTimeout time.Duration
|
||||||
@@ -60,6 +84,10 @@ type Config struct {
|
|||||||
RateLimitPerMinute int64
|
RateLimitPerMinute int64
|
||||||
RateLimitPerHour int64
|
RateLimitPerHour int64
|
||||||
RateLimitPerDay int64
|
RateLimitPerDay int64
|
||||||
|
// RateLimitExemptPaths are the path prefixes whose requests the rate
|
||||||
|
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
|
||||||
|
// Each starts with /.
|
||||||
|
RateLimitExemptPaths []string
|
||||||
// DeniedCountries are the countries whose clients are refused
|
// DeniedCountries are the countries whose clients are refused
|
||||||
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
|
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
|
||||||
// empty, are the only countries whose clients are let through
|
// empty, are the only countries whose clients are let through
|
||||||
@@ -67,9 +95,68 @@ type Config struct {
|
|||||||
// capitals, as GeoJS gives them.
|
// capitals, as GeoJS gives them.
|
||||||
DeniedCountries []string
|
DeniedCountries []string
|
||||||
ExclusivelyAllowedCountries []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 or matches a ban rule, 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
|
||||||
|
// AttackBanDuration is the ban for a first clear sign of attack
|
||||||
|
// (SWWAF_ATTACK_BAN_DURATION). It cannot be off.
|
||||||
|
AttackBanDuration 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
|
||||||
|
// LogRequestHeaders are the request headers whose values the request
|
||||||
|
// log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS).
|
||||||
|
LogRequestHeaders []string
|
||||||
|
// MetricsToken is the bearer token a scraper sends for the metrics
|
||||||
|
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
|
||||||
|
// MetricsTopN is how many countries get series of their own in the
|
||||||
|
// metrics (SWWAF_METRICS_TOP_N).
|
||||||
|
MetricsToken string
|
||||||
|
MetricsTopN int
|
||||||
|
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
|
||||||
|
// unless RulesEnabled is false (SWWAF_RULES_ENABLED).
|
||||||
|
RulesDir string
|
||||||
|
RulesEnabled bool
|
||||||
|
// LogRemoteURL is where every line on stdout is also sent
|
||||||
|
// (SWWAF_LOG_REMOTE_URL), nil while it is unset and nothing is sent.
|
||||||
|
// LogRemoteTLSCAs are the certificates a syslog+tls endpoint's
|
||||||
|
// certificate must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil while
|
||||||
|
// it is unset and the host's own are used. LogRemoteBuffer is the most
|
||||||
|
// lines held while they wait to be sent (SWWAF_LOG_REMOTE_BUFFER).
|
||||||
|
// LogRemoteFacility is the number of the syslog facility
|
||||||
|
// (SWWAF_LOG_REMOTE_FACILITY), and LogRemoteAppName the APP-NAME
|
||||||
|
// (SWWAF_LOG_REMOTE_APP_NAME, by default InstanceName), of the records
|
||||||
|
// the lines are sent in.
|
||||||
|
LogRemoteURL *url.URL
|
||||||
|
LogRemoteTLSCAs *x509.CertPool
|
||||||
|
LogRemoteBuffer int
|
||||||
|
LogRemoteFacility int
|
||||||
|
LogRemoteAppName string
|
||||||
|
|
||||||
// settings are the values read, as given or by default, for the
|
// settings are the values read, as given or by default, and the
|
||||||
// log line at start.
|
// files they were read from, for the log line at start.
|
||||||
settings []slog.Attr
|
settings []slog.Attr
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -82,6 +169,15 @@ const (
|
|||||||
kibibyte = 1 << 10
|
kibibyte = 1 << 10
|
||||||
mebibyte = 1 << 20
|
mebibyte = 1 << 20
|
||||||
gibibyte = 1 << 30
|
gibibyte = 1 << 30
|
||||||
|
ipv4Bits = 32
|
||||||
|
// minTokenLength is the fewest characters a token may have.
|
||||||
|
minTokenLength = 32
|
||||||
|
// masked is what the log shows for a token that is set.
|
||||||
|
masked = "********"
|
||||||
|
// defaultListenAddr and defaultUpstreamURL are the defaults of
|
||||||
|
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
|
||||||
|
defaultListenAddr = ":8080"
|
||||||
|
defaultUpstreamURL = "http://127.0.0.1:8081"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -102,19 +198,55 @@ var (
|
|||||||
"such as http://127.0.0.1:8081")
|
"such as http://127.0.0.1:8081")
|
||||||
errNotCountry = errors.New(
|
errNotCountry = errors.New(
|
||||||
"is not a two-letter country code such as de or kp")
|
"is not a two-letter country code such as de or kp")
|
||||||
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
|
errNotHeaderName = errors.New(
|
||||||
|
"is not a header name such as accept-language")
|
||||||
|
errHeaderTakenOut = errors.New(
|
||||||
|
"is taken out of every request by Go's HTTP server, so it can never " +
|
||||||
|
"be logged")
|
||||||
|
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
|
||||||
|
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")
|
||||||
|
errShortToken = errors.New("is shorter than 32 characters")
|
||||||
|
errNotMode = errors.New("is not enforce or observe")
|
||||||
|
errNotPathPrefix = errors.New(
|
||||||
|
"is not a path prefix starting with /, such as /assets/")
|
||||||
|
errNotBoolean = errors.New("is not true or false")
|
||||||
|
errNotLogRemoteURL = errors.New(
|
||||||
|
"is not syslog+udp, syslog+tcp or syslog+tls with a host and a port, " +
|
||||||
|
"and nothing more, such as syslog+tls://logs.example:6514")
|
||||||
|
errNoCertificate = errors.New("holds no PEM certificate")
|
||||||
|
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
|
||||||
|
errNotAppName = errors.New(
|
||||||
|
"is not 1 to 48 printable ASCII characters without a space, such as gitea")
|
||||||
|
errSetTwice = errors.New("set only one of them")
|
||||||
)
|
)
|
||||||
|
|
||||||
// FromEnvironment reads the settings with lookupEnv, normally
|
// FromEnvironment reads the settings with lookupEnv, normally
|
||||||
// os.LookupEnv. A setting that is not set takes its default. A setting
|
// os.LookupEnv. A setting may instead be given as a file: the variable
|
||||||
// that is set but invalid is an error that names it.
|
// named by the setting's name with _FILE added names the file, which is
|
||||||
|
// read now (see lookup). A setting that is not set takes its default. A
|
||||||
|
// setting that is set but invalid is an error that names it.
|
||||||
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||||
env := &environment{lookupEnv: lookupEnv}
|
env := &environment{lookupEnv: lookupEnv}
|
||||||
|
hostname, _ := os.Hostname() // "" when the host has no name to give
|
||||||
cfg := &Config{
|
cfg := &Config{
|
||||||
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
|
ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
|
||||||
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
|
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
|
||||||
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
|
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
|
||||||
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
|
Observe: env.observe("SWWAF_MODE", "enforce"),
|
||||||
|
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"),
|
ClientResponseTimeout: env.duration("SWWAF_CLIENT_RESPONSE_TIMEOUT", "30m"),
|
||||||
UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"),
|
UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"),
|
||||||
UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"),
|
UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"),
|
||||||
@@ -126,11 +258,35 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
|||||||
RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"),
|
RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"),
|
||||||
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
|
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
|
||||||
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
|
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
|
||||||
|
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
|
||||||
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
|
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
|
||||||
ExclusivelyAllowedCountries: env.countries(
|
ExclusivelyAllowedCountries: env.countries(
|
||||||
"SWWAF_EXCLUSIVELY_ALLOWED_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"),
|
||||||
|
AttackBanDuration: env.durationNotOff("SWWAF_ATTACK_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"),
|
||||||
|
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
|
||||||
|
"accept,accept-language,accept-encoding,content-type,origin,range"),
|
||||||
|
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
|
||||||
|
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
|
||||||
|
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
|
||||||
|
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
|
||||||
|
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
|
||||||
|
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
|
||||||
|
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
|
||||||
|
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
|
||||||
|
cfg.InstanceName, cfg.LogRemoteURL != nil)
|
||||||
|
|
||||||
for _, country := range cfg.ExclusivelyAllowedCountries {
|
for _, country := range cfg.ExclusivelyAllowedCountries {
|
||||||
if slices.Contains(cfg.DeniedCountries, country) {
|
if slices.Contains(cfg.DeniedCountries, country) {
|
||||||
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
|
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
|
||||||
@@ -147,6 +303,24 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
|||||||
return cfg, nil
|
return cfg, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ListenAddrAndUpstreamURL reads only SWWAF_LISTEN_ADDR and
|
||||||
|
// SWWAF_UPSTREAM_URL, either of which may be given as a file, as
|
||||||
|
// FromEnvironment does. The health check needs no other setting, so it
|
||||||
|
// reads no other, nor a file that another names.
|
||||||
|
func ListenAddrAndUpstreamURL(
|
||||||
|
lookupEnv func(string) (string, bool),
|
||||||
|
) (string, *url.URL, error) {
|
||||||
|
env := &environment{lookupEnv: lookupEnv}
|
||||||
|
listenAddr := env.address("SWWAF_LISTEN_ADDR", defaultListenAddr)
|
||||||
|
upstreamURL := env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL)
|
||||||
|
|
||||||
|
if env.err != nil {
|
||||||
|
return "", nil, env.err
|
||||||
|
}
|
||||||
|
|
||||||
|
return listenAddr, upstreamURL, nil
|
||||||
|
}
|
||||||
|
|
||||||
// privateRanges are the private address ranges, the default trusted
|
// privateRanges are the private address ranges, the default trusted
|
||||||
// proxies.
|
// proxies.
|
||||||
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
|
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
|
||||||
@@ -168,8 +342,8 @@ type environment struct {
|
|||||||
// value returns a setting's value, or its default when it is not set,
|
// value returns a setting's value, or its default when it is not set,
|
||||||
// and notes it for the log.
|
// and notes it for the log.
|
||||||
func (e *environment) value(name, defaultValue string) string {
|
func (e *environment) value(name, defaultValue string) string {
|
||||||
value, ok := e.lookupEnv(name)
|
value, set := e.lookup(name)
|
||||||
if !ok {
|
if !set {
|
||||||
value = defaultValue
|
value = defaultValue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -178,6 +352,37 @@ func (e *environment) value(name, defaultValue string) string {
|
|||||||
return value
|
return value
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// lookup returns a setting's value and whether it is set: the value of the
|
||||||
|
// variable name, or the contents of the file that the variable name_FILE
|
||||||
|
// names, less one newline at their end. It notes that file's path for the
|
||||||
|
// log. Both variables set, or a file that cannot be read, is an error.
|
||||||
|
func (e *environment) lookup(name string) (string, bool) {
|
||||||
|
value, set := e.lookupEnv(name)
|
||||||
|
fileName := name + "_FILE"
|
||||||
|
|
||||||
|
path, inFile := e.lookupEnv(fileName)
|
||||||
|
if !inFile {
|
||||||
|
return value, set
|
||||||
|
}
|
||||||
|
|
||||||
|
if set {
|
||||||
|
e.check(name, fmt.Errorf("is set, and so is %s; %w", fileName, errSetTwice))
|
||||||
|
|
||||||
|
return value, set
|
||||||
|
}
|
||||||
|
|
||||||
|
e.settings = append(e.settings, slog.String(fileName, path))
|
||||||
|
|
||||||
|
contents, err := os.ReadFile(path) //nolint:gosec // a file the admin names
|
||||||
|
if err != nil {
|
||||||
|
e.check(fileName, fmt.Errorf("cannot be read: %w", err))
|
||||||
|
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimSuffix(string(contents), "\n"), true
|
||||||
|
}
|
||||||
|
|
||||||
// check keeps the first error, naming the setting it is about.
|
// check keeps the first error, naming the setting it is about.
|
||||||
func (e *environment) check(name string, err error) {
|
func (e *environment) check(name string, err error) {
|
||||||
if err != nil && e.err == nil {
|
if err != nil && e.err == nil {
|
||||||
@@ -201,6 +406,27 @@ func (e *environment) appURL(name, defaultValue string) *url.URL {
|
|||||||
return upstream
|
return upstream
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// observe reads the setting that is the mode, enforce or observe, and
|
||||||
|
// reports whether it is observe.
|
||||||
|
func (e *environment) observe(name, defaultValue string) bool {
|
||||||
|
mode := e.value(name, defaultValue)
|
||||||
|
if mode != "enforce" && mode != "observe" {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", mode, errNotMode))
|
||||||
|
}
|
||||||
|
|
||||||
|
return mode == "observe"
|
||||||
|
}
|
||||||
|
|
||||||
|
// boolean reads a setting that is true or false.
|
||||||
|
func (e *environment) boolean(name, defaultValue string) bool {
|
||||||
|
value := e.value(name, defaultValue)
|
||||||
|
if value != "true" && value != "false" {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", value, errNotBoolean))
|
||||||
|
}
|
||||||
|
|
||||||
|
return value == "true"
|
||||||
|
}
|
||||||
|
|
||||||
// netblocks reads a setting that is a list of netblocks.
|
// netblocks reads a setting that is a list of netblocks.
|
||||||
func (e *environment) netblocks(name, defaultValue string) []netip.Prefix {
|
func (e *environment) netblocks(name, defaultValue string) []netip.Prefix {
|
||||||
netblocks, err := parseNetblocks(e.value(name, defaultValue))
|
netblocks, err := parseNetblocks(e.value(name, defaultValue))
|
||||||
@@ -225,6 +451,15 @@ func (e *environment) size(name, defaultValue string) int64 {
|
|||||||
return size
|
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.
|
// count reads a setting that is a number of requests.
|
||||||
func (e *environment) count(name, defaultValue string) int64 {
|
func (e *environment) count(name, defaultValue string) int64 {
|
||||||
count, err := parseCount(e.value(name, defaultValue))
|
count, err := parseCount(e.value(name, defaultValue))
|
||||||
@@ -233,6 +468,14 @@ func (e *environment) count(name, defaultValue string) int64 {
|
|||||||
return count
|
return count
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// pathPrefixes reads a setting that is a list of path prefixes.
|
||||||
|
func (e *environment) pathPrefixes(name, defaultValue string) []string {
|
||||||
|
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return prefixes
|
||||||
|
}
|
||||||
|
|
||||||
// countries reads a setting that is a list of countries.
|
// countries reads a setting that is a list of countries.
|
||||||
func (e *environment) countries(name, defaultValue string) []string {
|
func (e *environment) countries(name, defaultValue string) []string {
|
||||||
countries, err := parseCountries(e.value(name, defaultValue))
|
countries, err := parseCountries(e.value(name, defaultValue))
|
||||||
@@ -241,6 +484,153 @@ func (e *environment) countries(name, defaultValue string) []string {
|
|||||||
return countries
|
return countries
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// headerNames reads a setting that is a list of header names, and
|
||||||
|
// returns them in lower case.
|
||||||
|
func (e *environment) headerNames(name, defaultValue string) []string {
|
||||||
|
headers, err := parseHeaderNames(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
|
// durationNotOff reads a setting that is a duration and, unlike a
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// token reads a setting that is a bearer token. Unset, it is "", which
|
||||||
|
// switches off what it guards; set, it must be at least minTokenLength
|
||||||
|
// characters. Neither the log nor an error shows its value.
|
||||||
|
func (e *environment) token(name string) string {
|
||||||
|
value, set := e.lookup(name)
|
||||||
|
if !set {
|
||||||
|
e.settings = append(e.settings, slog.String(name, ""))
|
||||||
|
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
e.settings = append(e.settings, slog.String(name, masked))
|
||||||
|
|
||||||
|
if utf8.RuneCountInString(value) < minTokenLength {
|
||||||
|
e.check(name, errShortToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// logRemoteURL reads the setting that is where every log line is also
|
||||||
|
// sent. Unset or empty, it is nil, and nothing is sent.
|
||||||
|
func (e *environment) logRemoteURL(name string) *url.URL {
|
||||||
|
value := e.value(name, "")
|
||||||
|
if value == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
remote, err := parseLogRemoteURL(value)
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return remote
|
||||||
|
}
|
||||||
|
|
||||||
|
// certificates reads a setting that is the path of a file of PEM
|
||||||
|
// certificates. Unset or empty, it is nil. Its value names a file
|
||||||
|
// already, so, unlike the other settings, it has no _FILE form.
|
||||||
|
func (e *environment) certificates(name string) *x509.CertPool {
|
||||||
|
path, _ := e.lookupEnv(name)
|
||||||
|
e.settings = append(e.settings, slog.String(name, path))
|
||||||
|
|
||||||
|
if path == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
pem, err := os.ReadFile(path) //nolint:gosec // a file the admin names
|
||||||
|
if err != nil {
|
||||||
|
e.check(name, fmt.Errorf("cannot be read: %w", err))
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
pool := x509.NewCertPool()
|
||||||
|
if !pool.AppendCertsFromPEM(pem) {
|
||||||
|
e.check(name, fmt.Errorf("%q %w", path, errNoCertificate))
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return pool
|
||||||
|
}
|
||||||
|
|
||||||
|
// facility reads a setting that is a syslog facility, and returns its
|
||||||
|
// number.
|
||||||
|
func (e *environment) facility(name, defaultValue string) int {
|
||||||
|
number, err := parseFacility(e.value(name, defaultValue))
|
||||||
|
e.check(name, err)
|
||||||
|
|
||||||
|
return number
|
||||||
|
}
|
||||||
|
|
||||||
|
// appName reads the setting that is the APP-NAME of the records the log
|
||||||
|
// lines are sent in, by default the instance name. Its value is checked
|
||||||
|
// when it is set, and, while lines are sent, when it is the instance name.
|
||||||
|
func (e *environment) appName(name, instanceName string, sending bool) string {
|
||||||
|
value, set := e.lookup(name)
|
||||||
|
if !set {
|
||||||
|
value = instanceName
|
||||||
|
}
|
||||||
|
|
||||||
|
e.settings = append(e.settings, slog.String(name, value))
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case isAppName(value):
|
||||||
|
case set:
|
||||||
|
e.check(name, fmt.Errorf("%q %w", value, errNotAppName))
|
||||||
|
case sending:
|
||||||
|
e.check(name, fmt.Errorf("is unset, and SWWAF_INSTANCE_NAME %q, its default, %w",
|
||||||
|
value, errNotAppName))
|
||||||
|
}
|
||||||
|
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
||||||
// whole number of days such as 7d, or off.
|
// whole number of days such as 7d, or off.
|
||||||
func parseDuration(value string) (time.Duration, error) {
|
func parseDuration(value string) (time.Duration, error) {
|
||||||
@@ -296,6 +686,19 @@ func parseSize(value string) (int64, error) {
|
|||||||
return n * unit, nil
|
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
|
// splitUnit splits a size into its number and the bytes its suffix
|
||||||
// stands for.
|
// stands for.
|
||||||
func splitUnit(value string) (string, int64) {
|
func splitUnit(value string) (string, int64) {
|
||||||
@@ -329,6 +732,52 @@ func parseCount(value string) (int64, error) {
|
|||||||
return n, nil
|
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
|
// parseList splits a comma-separated list and trims the spaces around
|
||||||
// each item. An empty value is an empty list.
|
// each item. An empty value is an empty list.
|
||||||
func parseList(value string) ([]string, error) {
|
func parseList(value string) ([]string, error) {
|
||||||
@@ -388,6 +837,23 @@ func parseNetblock(value string) (netip.Prefix, error) {
|
|||||||
return netip.PrefixFrom(addr, addr.BitLen()), nil
|
return netip.PrefixFrom(addr, addr.BitLen()), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parsePathPrefixes reads a comma-separated list of path prefixes, each
|
||||||
|
// starting with /.
|
||||||
|
func parsePathPrefixes(value string) ([]string, error) {
|
||||||
|
prefixes, err := parseList(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, prefix := range prefixes {
|
||||||
|
if !strings.HasPrefix(prefix, "/") {
|
||||||
|
return nil, fmt.Errorf("%q %w", prefix, errNotPathPrefix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return prefixes, nil
|
||||||
|
}
|
||||||
|
|
||||||
// countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK,
|
// countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK,
|
||||||
// the code in common use for Kosovo. golang.org/x/text/language cannot
|
// the code in common use for Kosovo. golang.org/x/text/language cannot
|
||||||
// check them: it also takes withdrawn codes such as su, and reserved ones
|
// check them: it also takes withdrawn codes such as su, and reserved ones
|
||||||
@@ -444,6 +910,58 @@ func parseCountries(value string) ([]string, error) {
|
|||||||
return countries, nil
|
return countries, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// headerNameChars are the characters RFC 9110 allows in a header name:
|
||||||
|
// letters, digits and these marks.
|
||||||
|
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
|
||||||
|
"0123456789!#$%&'*+-.^_`|~"
|
||||||
|
|
||||||
|
// IsHeaderName reports whether name can be a header name: one or more of
|
||||||
|
// the characters RFC 9110 allows in one.
|
||||||
|
func IsHeaderName(name string) bool {
|
||||||
|
if name == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, char := range name {
|
||||||
|
if !strings.ContainsRune(headerNameChars, char) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseHeaderNames reads a comma-separated list of header names in either
|
||||||
|
// case, and returns them in lower case. Host and Transfer-Encoding are
|
||||||
|
// refused: Go's HTTP server takes them out of the request's headers.
|
||||||
|
func parseHeaderNames(value string) ([]string, error) {
|
||||||
|
items, err := parseList(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
headers := make([]string, 0, len(items))
|
||||||
|
|
||||||
|
for _, item := range items {
|
||||||
|
if !IsHeaderName(item) {
|
||||||
|
return nil, fmt.Errorf("%q %w", item, errNotHeaderName)
|
||||||
|
}
|
||||||
|
|
||||||
|
header := strings.ToLower(item)
|
||||||
|
switch header {
|
||||||
|
case "host":
|
||||||
|
return nil, fmt.Errorf("%q %w; the request's host is the field host",
|
||||||
|
item, errHeaderTakenOut)
|
||||||
|
case "transfer-encoding":
|
||||||
|
return nil, fmt.Errorf("%q %w", item, errHeaderTakenOut)
|
||||||
|
}
|
||||||
|
|
||||||
|
headers = append(headers, header)
|
||||||
|
}
|
||||||
|
|
||||||
|
return headers, nil
|
||||||
|
}
|
||||||
|
|
||||||
// parseListenAddr checks an address to listen on: an optional host and a
|
// parseListenAddr checks an address to listen on: an optional host and a
|
||||||
// port number.
|
// port number.
|
||||||
func parseListenAddr(value string) (string, error) {
|
func parseListenAddr(value string) (string, error) {
|
||||||
@@ -486,3 +1004,64 @@ func parseUpstreamURL(value string) (*url.URL, error) {
|
|||||||
|
|
||||||
return upstream, nil
|
return upstream, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseLogRemoteURL reads where every log line is also sent:
|
||||||
|
// syslog+udp, syslog+tcp or syslog+tls, a host and a port from 1 to
|
||||||
|
// 65535, and nothing else.
|
||||||
|
func parseLogRemoteURL(value string) (*url.URL, error) {
|
||||||
|
remote, err := url.Parse(value)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%q %w", value, errNotLogRemoteURL)
|
||||||
|
}
|
||||||
|
|
||||||
|
schemes := []string{remotelog.SchemeUDP, remotelog.SchemeTCP, remotelog.SchemeTLS}
|
||||||
|
port, err := strconv.ParseUint(remote.Port(), 10, 16)
|
||||||
|
|
||||||
|
onlySchemeHostAndPort := slices.Contains(schemes, remote.Scheme) &&
|
||||||
|
remote.Hostname() != "" && err == nil && port != 0 &&
|
||||||
|
remote.User == nil && remote.Opaque == "" &&
|
||||||
|
(remote.Path == "" || remote.Path == "/") &&
|
||||||
|
remote.RawQuery == "" && remote.Fragment == ""
|
||||||
|
if !onlySchemeHostAndPort {
|
||||||
|
return nil, fmt.Errorf("%q %w", value, errNotLogRemoteURL)
|
||||||
|
}
|
||||||
|
|
||||||
|
return remote, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseFacility reads the name of a syslog facility, and returns its
|
||||||
|
// number, as RFC 5424 numbers them.
|
||||||
|
func parseFacility(value string) (int, error) {
|
||||||
|
//nolint:mnd // the facilities' numbers in RFC 5424
|
||||||
|
number, known := map[string]int{
|
||||||
|
"kern": 0, "user": 1, "mail": 2, "daemon": 3, "auth": 4, "syslog": 5,
|
||||||
|
"lpr": 6, "news": 7, "uucp": 8, "cron": 9, "authpriv": 10, "ftp": 11,
|
||||||
|
"local0": 16, "local1": 17, "local2": 18, "local3": 19,
|
||||||
|
"local4": 20, "local5": 21, "local6": 22, "local7": 23,
|
||||||
|
}[value]
|
||||||
|
if !known {
|
||||||
|
return 0, fmt.Errorf("%q %w", value, errNotFacility)
|
||||||
|
}
|
||||||
|
|
||||||
|
return number, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// appNameMaxLength is the most characters RFC 5424 allows in an
|
||||||
|
// APP-NAME.
|
||||||
|
const appNameMaxLength = 48
|
||||||
|
|
||||||
|
// isAppName reports whether value can be an APP-NAME: 1 to
|
||||||
|
// appNameMaxLength printable ASCII characters, none of them a space.
|
||||||
|
func isAppName(value string) bool {
|
||||||
|
if value == "" || len(value) > appNameMaxLength {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, char := range []byte(value) {
|
||||||
|
if char < '!' || char > '~' {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|||||||
+679
-32
@@ -2,10 +2,13 @@ package config_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"crypto/x509"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"maps"
|
"maps"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -18,8 +21,11 @@ import (
|
|||||||
const (
|
const (
|
||||||
listenAddr = "SWWAF_LISTEN_ADDR"
|
listenAddr = "SWWAF_LISTEN_ADDR"
|
||||||
upstreamURL = "SWWAF_UPSTREAM_URL"
|
upstreamURL = "SWWAF_UPSTREAM_URL"
|
||||||
|
mode = "SWWAF_MODE"
|
||||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||||
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
||||||
|
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
|
||||||
|
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
|
||||||
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
||||||
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
||||||
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
||||||
@@ -31,8 +37,58 @@ const (
|
|||||||
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
||||||
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
|
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
|
||||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
|
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
||||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||||
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_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"
|
||||||
|
attackBanDuration = "SWWAF_ATTACK_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"
|
||||||
|
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
|
||||||
|
metricsTopN = "SWWAF_METRICS_TOP_N"
|
||||||
|
instanceName = "SWWAF_INSTANCE_NAME"
|
||||||
|
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
|
||||||
|
rulesDir = "SWWAF_RULES_DIR"
|
||||||
|
rulesEnabled = "SWWAF_RULES_ENABLED"
|
||||||
|
logRemoteURL = "SWWAF_LOG_REMOTE_URL"
|
||||||
|
logRemoteTLSCAFile = "SWWAF_LOG_REMOTE_TLS_CA_FILE"
|
||||||
|
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
|
||||||
|
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
|
||||||
|
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
|
||||||
|
)
|
||||||
|
|
||||||
|
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
|
||||||
|
const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," +
|
||||||
|
"content-type,origin,range"
|
||||||
|
|
||||||
|
// testCA is a CA certificate, of which only that it reads matters here.
|
||||||
|
const testCA = `-----BEGIN CERTIFICATE-----
|
||||||
|
MIIBkzCCATmgAwIBAgIUeySaE27dnr6A2HijrMB13gTLUKIwCgYIKoZIzj0EAwIw
|
||||||
|
HjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBDQTAgFw0yNjEwMDYxNDI2MTNa
|
||||||
|
GA8yMTI2MDkxMjE0MjYxM1owHjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBD
|
||||||
|
QTBZMBMGByqGSM49AgEGCCqGSM49AwEHA0IABKDEhcWKKhet2KgSdME+iEPxyEyn
|
||||||
|
2sd9IdElbt8DM2SfCdB2JsXo0C07UNZaywMPMfn/n8LNI/PKwu+N2uX7gfSjUzBR
|
||||||
|
MB0GA1UdDgQWBBRZo3BPLv0KbV4drw6JI1JIUKmoRzAfBgNVHSMEGDAWgBRZo3BP
|
||||||
|
Lv0KbV4drw6JI1JIUKmoRzAPBgNVHRMBAf8EBTADAQH/MAoGCCqGSM49BAMCA0gA
|
||||||
|
MEUCICC5k+76UpWoSwVbZA+atu5WcALEOJGqwOUWua3zemhcAiEA8Hfdxgwp0z2v
|
||||||
|
rlG9y/jrJb6ORy3kTLWo2EA0BA67vuI=
|
||||||
|
-----END CERTIFICATE-----
|
||||||
|
`
|
||||||
|
|
||||||
|
// token is a token of 32 characters, the shortest allowed.
|
||||||
|
const token = "0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
// instance is an SWWAF_INSTANCE_NAME that is a valid app name too, and
|
||||||
|
// remoteURL an SWWAF_LOG_REMOTE_URL, for the tests that send the lines.
|
||||||
|
const (
|
||||||
|
instance = "fsn1app1/gitea"
|
||||||
|
remoteURL = "syslog+udp://192.0.2.1:514"
|
||||||
)
|
)
|
||||||
|
|
||||||
// off switches a timeout, a size limit or a rate limit off.
|
// off switches a timeout, a size limit or a rate limit off.
|
||||||
@@ -66,16 +122,33 @@ func TestDefaults(t *testing.T) {
|
|||||||
cfg := fromEnvironment(t, environment{})
|
cfg := fromEnvironment(t, environment{})
|
||||||
|
|
||||||
wantSettings(t, cfg, config.Config{
|
wantSettings(t, cfg, config.Config{
|
||||||
ListenAddr: ":8080",
|
ListenAddr: ":8080",
|
||||||
ClientRequestTimeout: time.Minute,
|
Observe: false,
|
||||||
ClientResponseTimeout: 30 * time.Minute,
|
ClientRequestTimeout: time.Minute,
|
||||||
UpstreamRequestTimeout: time.Minute,
|
ClientRequestHeaderMaxBytes: 32 << 10,
|
||||||
UpstreamResponseTimeout: 30 * time.Minute,
|
ClientIdleTimeout: 2 * time.Minute,
|
||||||
RequestMaxBytes: 100 << 20,
|
ClientResponseTimeout: 30 * time.Minute,
|
||||||
ResponseMaxBytes: 5 << 30,
|
UpstreamRequestTimeout: time.Minute,
|
||||||
RateLimitPerMinute: 1000,
|
UpstreamResponseTimeout: 30 * time.Minute,
|
||||||
RateLimitPerHour: 10000,
|
RequestMaxBytes: 100 << 20,
|
||||||
RateLimitPerDay: 50000,
|
ResponseMaxBytes: 5 << 30,
|
||||||
|
RateLimitPerMinute: 1000,
|
||||||
|
RateLimitPerHour: 10000,
|
||||||
|
RateLimitPerDay: 50000,
|
||||||
|
BanResponse: 403,
|
||||||
|
LimitBanDuration: time.Hour,
|
||||||
|
LimitBanRepeatWindow: 24 * time.Hour,
|
||||||
|
MaxBanDuration: 7 * 24 * time.Hour,
|
||||||
|
AttackBanDuration: 7 * 24 * time.Hour,
|
||||||
|
MaxBans: 5000,
|
||||||
|
BanScopeV4Prefix: 32,
|
||||||
|
StateDir: "/var/lib/smallwebwaf",
|
||||||
|
StateWriteDelay: 10 * time.Second,
|
||||||
|
StateCounterInterval: 15 * time.Minute,
|
||||||
|
MetricsToken: "",
|
||||||
|
MetricsTopN: 50,
|
||||||
|
RulesDir: "/etc/smallwebwaf/rules.d",
|
||||||
|
RulesEnabled: true,
|
||||||
})
|
})
|
||||||
|
|
||||||
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
||||||
@@ -89,6 +162,22 @@ func TestDefaults(t *testing.T) {
|
|||||||
wantNetblocks(t, cfg.DenyNets)
|
wantNetblocks(t, cfg.DenyNets)
|
||||||
wantCountries(t, deniedCountries, cfg.DeniedCountries)
|
wantCountries(t, deniedCountries, cfg.DeniedCountries)
|
||||||
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries)
|
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries)
|
||||||
|
|
||||||
|
hostname, err := os.Hostname()
|
||||||
|
if err != nil || hostname == "" || cfg.InstanceName != hostname {
|
||||||
|
t.Errorf("%s is %q, want the host's name %q (%v)", instanceName,
|
||||||
|
cfg.InstanceName, hostname, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantHeaders := strings.Split(defaultLogRequestHeaders, ",")
|
||||||
|
if !slices.Equal(cfg.LogRequestHeaders, wantHeaders) {
|
||||||
|
t.Errorf("%s gave %v, want %v", logRequestHeaders, cfg.LogRequestHeaders,
|
||||||
|
wantHeaders)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.RateLimitExemptPaths) != 0 {
|
||||||
|
t.Errorf("%s gave %v, want none", rateLimitExemptPaths, cfg.RateLimitExemptPaths)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValuesAsSet(t *testing.T) {
|
func TestValuesAsSet(t *testing.T) {
|
||||||
@@ -97,8 +186,11 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
cfg := fromEnvironment(t, environment{
|
cfg := fromEnvironment(t, environment{
|
||||||
listenAddr: "127.0.0.1:9000",
|
listenAddr: "127.0.0.1:9000",
|
||||||
upstreamURL: "https://app.internal:8443/",
|
upstreamURL: "https://app.internal:8443/",
|
||||||
|
mode: "observe",
|
||||||
trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32",
|
trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32",
|
||||||
clientRequestTimeout: "90s",
|
clientRequestTimeout: "90s",
|
||||||
|
clientHeaderMaxBytes: "8K",
|
||||||
|
clientIdleTimeout: "5m",
|
||||||
clientResponseTimeout: "7d",
|
clientResponseTimeout: "7d",
|
||||||
upstreamRequestTimeout: "1h30m",
|
upstreamRequestTimeout: "1h30m",
|
||||||
upstreamResponseTimeout: off,
|
upstreamResponseTimeout: off,
|
||||||
@@ -112,19 +204,50 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
rateLimitPerDay: "6000",
|
rateLimitPerDay: "6000",
|
||||||
deniedCountries: "cn, RU,kp,Xk",
|
deniedCountries: "cn, RU,kp,Xk",
|
||||||
allowedCountries: "de",
|
allowedCountries: "de",
|
||||||
|
banResponse: "429",
|
||||||
|
limitBanDuration: "15m",
|
||||||
|
limitBanRepeatWindow: "2d",
|
||||||
|
maxBanDuration: "30d",
|
||||||
|
attackBanDuration: "1d",
|
||||||
|
maxBans: "100",
|
||||||
|
banScopeV4Prefix: "24",
|
||||||
|
stateDir: "/srv/waf-state",
|
||||||
|
stateWriteDelay: "500ms",
|
||||||
|
stateCounterInterval: "1h",
|
||||||
|
metricsToken: token,
|
||||||
|
metricsTopN: "10",
|
||||||
|
rulesDir: "/srv/waf-rules",
|
||||||
|
rulesEnabled: "false",
|
||||||
})
|
})
|
||||||
|
|
||||||
wantSettings(t, cfg, config.Config{
|
wantSettings(t, cfg, config.Config{
|
||||||
ListenAddr: "127.0.0.1:9000",
|
ListenAddr: "127.0.0.1:9000",
|
||||||
ClientRequestTimeout: 90 * time.Second,
|
Observe: true,
|
||||||
ClientResponseTimeout: 7 * 24 * time.Hour,
|
ClientRequestTimeout: 90 * time.Second,
|
||||||
UpstreamRequestTimeout: 90 * time.Minute,
|
ClientRequestHeaderMaxBytes: 8 << 10,
|
||||||
UpstreamResponseTimeout: 0,
|
ClientIdleTimeout: 5 * time.Minute,
|
||||||
RequestMaxBytes: 512 << 10,
|
ClientResponseTimeout: 7 * 24 * time.Hour,
|
||||||
ResponseMaxBytes: 1234,
|
UpstreamRequestTimeout: 90 * time.Minute,
|
||||||
RateLimitPerMinute: 60,
|
UpstreamResponseTimeout: 0,
|
||||||
RateLimitPerHour: 600,
|
RequestMaxBytes: 512 << 10,
|
||||||
RateLimitPerDay: 6000,
|
ResponseMaxBytes: 1234,
|
||||||
|
RateLimitPerMinute: 60,
|
||||||
|
RateLimitPerHour: 600,
|
||||||
|
RateLimitPerDay: 6000,
|
||||||
|
BanResponse: 429,
|
||||||
|
LimitBanDuration: 15 * time.Minute,
|
||||||
|
LimitBanRepeatWindow: 48 * time.Hour,
|
||||||
|
MaxBanDuration: 30 * 24 * time.Hour,
|
||||||
|
AttackBanDuration: 24 * time.Hour,
|
||||||
|
MaxBans: 100,
|
||||||
|
BanScopeV4Prefix: 24,
|
||||||
|
StateDir: "/srv/waf-state",
|
||||||
|
StateWriteDelay: 500 * time.Millisecond,
|
||||||
|
StateCounterInterval: time.Hour,
|
||||||
|
MetricsToken: token,
|
||||||
|
MetricsTopN: 10,
|
||||||
|
RulesDir: "/srv/waf-rules",
|
||||||
|
RulesEnabled: false,
|
||||||
})
|
})
|
||||||
|
|
||||||
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
|
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
|
||||||
@@ -139,6 +262,216 @@ func TestValuesAsSet(t *testing.T) {
|
|||||||
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
|
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRateLimitExemptPathsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{rateLimitExemptPaths: "/assets/, /favicon.ico"})
|
||||||
|
|
||||||
|
if !slices.Equal(cfg.RateLimitExemptPaths, []string{"/assets/", "/favicon.ico"}) {
|
||||||
|
t.Errorf("%s gave %v, want /assets/ and /favicon.ico",
|
||||||
|
rateLimitExemptPaths, cfg.RateLimitExemptPaths)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPathPrefixNotStartingWithSlashStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(
|
||||||
|
environment{rateLimitExemptPaths: "/favicon.ico,assets/"}.lookupEnv)
|
||||||
|
|
||||||
|
want := rateLimitExemptPaths + `: "assets/" is not a path prefix ` +
|
||||||
|
`starting with /, such as /assets/`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
instanceName: "fsn1app1/gitea",
|
||||||
|
logRequestHeaders: " Accept , X-Custom",
|
||||||
|
})
|
||||||
|
|
||||||
|
if cfg.InstanceName != "fsn1app1/gitea" ||
|
||||||
|
!slices.Equal(cfg.LogRequestHeaders, []string{"accept", "x-custom"}) {
|
||||||
|
t.Errorf("%s is %q and %s gives %v", instanceName, cfg.InstanceName,
|
||||||
|
logRequestHeaders, cfg.LogRequestHeaders)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogSettingsDefaults(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{instanceName: instance})
|
||||||
|
|
||||||
|
if cfg.LogRemoteURL != nil || cfg.LogRemoteTLSCAs != nil ||
|
||||||
|
cfg.LogRemoteBuffer != 10000 || cfg.LogRemoteFacility != 16 ||
|
||||||
|
cfg.LogRemoteAppName != instance {
|
||||||
|
t.Errorf("remote log settings %v, %v, %d, %d and %q, want no URL, no "+
|
||||||
|
"certificates, 10000, 16 and %s's %s", cfg.LogRemoteURL,
|
||||||
|
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
|
||||||
|
cfg.LogRemoteAppName, instanceName, instance)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogSettingsAsSet(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
caFile := filepath.Join(t.TempDir(), "ca.pem")
|
||||||
|
|
||||||
|
err := os.WriteFile(caFile, []byte(testCA), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", caFile, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
logRemoteURL: "syslog+tls://logs.example:6514",
|
||||||
|
logRemoteTLSCAFile: caFile,
|
||||||
|
logRemoteBuffer: "500",
|
||||||
|
logRemoteFacility: "daemon",
|
||||||
|
logRemoteAppName: instance,
|
||||||
|
})
|
||||||
|
|
||||||
|
roots := x509.NewCertPool()
|
||||||
|
roots.AppendCertsFromPEM([]byte(testCA))
|
||||||
|
|
||||||
|
if cfg.LogRemoteURL.String() != "syslog+tls://logs.example:6514" ||
|
||||||
|
!roots.Equal(cfg.LogRemoteTLSCAs) || cfg.LogRemoteBuffer != 500 ||
|
||||||
|
cfg.LogRemoteFacility != 3 || cfg.LogRemoteAppName != instance {
|
||||||
|
t.Errorf("remote log settings %v, %v, %d, %d and %q", cfg.LogRemoteURL,
|
||||||
|
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
|
||||||
|
cfg.LogRemoteAppName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogURLForms(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, value := range []string{
|
||||||
|
"syslog+udp://192.0.2.1:514",
|
||||||
|
"syslog+tcp://[2001:db8::1]:514",
|
||||||
|
"syslog+tls://logs.example:6514/",
|
||||||
|
} {
|
||||||
|
cfg := fromEnvironment(t, environment{logRemoteURL: value})
|
||||||
|
if cfg.LogRemoteURL.String() != value {
|
||||||
|
t.Errorf("%s read as %v", value, cfg.LogRemoteURL)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{logRemoteURL: ""})
|
||||||
|
if cfg.LogRemoteURL != nil {
|
||||||
|
t.Errorf("set but empty, %s read as %v", logRemoteURL, cfg.LogRemoteURL)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogFacilitiesByNumber(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for name, number := range map[string]int{
|
||||||
|
"kern": 0, "user": 1, "auth": 4, "authpriv": 10, "ftp": 11,
|
||||||
|
"local0": 16, "local5": 21, "local7": 23,
|
||||||
|
} {
|
||||||
|
cfg := fromEnvironment(t, environment{logRemoteFacility: name})
|
||||||
|
if cfg.LogRemoteFacility != number {
|
||||||
|
t.Errorf("%s read as %d, want %d", name, cfg.LogRemoteFacility, number)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInvalidRemoteLogSettingStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct{ name, value string }{
|
||||||
|
{logRemoteURL, "logs.example:514"},
|
||||||
|
{logRemoteURL, "syslog://logs.example:514"},
|
||||||
|
{logRemoteURL, "http://logs.example:514"},
|
||||||
|
{logRemoteURL, "syslog+udp://logs.example"},
|
||||||
|
{logRemoteURL, "syslog+tcp://:514"},
|
||||||
|
{logRemoteURL, "syslog+tcp://logs.example:0"},
|
||||||
|
{logRemoteURL, "syslog+tls://logs.example:65536"},
|
||||||
|
{logRemoteURL, "syslog+tls://user@logs.example:6514"},
|
||||||
|
{logRemoteURL, "syslog+tcp://logs.example:514/app"},
|
||||||
|
{logRemoteURL, "syslog+tcp://logs.example:514?tls=1"},
|
||||||
|
{logRemoteTLSCAFile, "/nonexistent/ca.pem"},
|
||||||
|
{logRemoteBuffer, off}, {logRemoteBuffer, "0"}, {logRemoteBuffer, "10K"},
|
||||||
|
{logRemoteFacility, "local8"}, {logRemoteFacility, "LOCAL0"},
|
||||||
|
{logRemoteFacility, "16"}, {logRemoteFacility, ""},
|
||||||
|
{logRemoteAppName, ""}, {logRemoteAppName, "my app"},
|
||||||
|
{logRemoteAppName, "gitéa"}, {logRemoteAppName, strings.Repeat("a", 49)},
|
||||||
|
} {
|
||||||
|
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
|
||||||
|
if err == nil || !strings.HasPrefix(err.Error(), tc.name+": ") {
|
||||||
|
t.Errorf("%s=%q: error %v, want one naming it", tc.name, tc.value, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogCAFileWithoutCertificateStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
caFile := filepath.Join(t.TempDir(), "ca.pem")
|
||||||
|
|
||||||
|
err := os.WriteFile(caFile, []byte("not a certificate\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", caFile, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = config.FromEnvironment(environment{logRemoteTLSCAFile: caFile}.lookupEnv)
|
||||||
|
|
||||||
|
want := logRemoteTLSCAFile + `: "` + caFile + `" holds no PEM certificate`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInstanceNameNotAnAppNameStopsTheStartOnlyWhileSending(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const spaced = "fsn1 app1"
|
||||||
|
|
||||||
|
sending := environment{logRemoteURL: remoteURL, instanceName: spaced}
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(sending.lookupEnv)
|
||||||
|
|
||||||
|
want := logRemoteAppName + `: is unset, and ` + instanceName +
|
||||||
|
` "fsn1 app1", its default, is not 1 to 48 printable ASCII characters ` +
|
||||||
|
`without a space, such as gitea`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{instanceName: spaced})
|
||||||
|
if cfg.LogRemoteAppName != spaced {
|
||||||
|
t.Errorf("not sending, %s is %q", logRemoteAppName, cfg.LogRemoteAppName)
|
||||||
|
}
|
||||||
|
|
||||||
|
sending[logRemoteAppName] = instance
|
||||||
|
|
||||||
|
cfg = fromEnvironment(t, sending)
|
||||||
|
if cfg.LogRemoteAppName != instance {
|
||||||
|
t.Errorf("set to %s, %s is %q", instance, logRemoteAppName,
|
||||||
|
cfg.LogRemoteAppName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAppNameSetStopsTheStartWhileSending(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
logRemoteURL: remoteURL,
|
||||||
|
instanceName: instance,
|
||||||
|
logRemoteAppName: "my app",
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want := logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
|
||||||
|
`characters without a space, such as gitea`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -163,12 +496,42 @@ func TestSizesAndOff(t *testing.T) {
|
|||||||
requestMaxBytes: "3G",
|
requestMaxBytes: "3G",
|
||||||
responseMaxBytes: off,
|
responseMaxBytes: off,
|
||||||
clientRequestTimeout: off,
|
clientRequestTimeout: off,
|
||||||
|
clientIdleTimeout: off,
|
||||||
})
|
})
|
||||||
|
|
||||||
if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 ||
|
if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 ||
|
||||||
cfg.ClientRequestTimeout != 0 {
|
cfg.ClientRequestTimeout != 0 || cfg.ClientIdleTimeout != 0 {
|
||||||
t.Errorf("3G, off and off read as %d, %d and %s",
|
t.Errorf("3G, off, off and off read as %d, %d, %s and %s",
|
||||||
cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout)
|
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 +551,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) {
|
func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -200,10 +572,8 @@ func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
|
|||||||
func TestInvalidValueStopsTheStart(t *testing.T) {
|
func TestInvalidValueStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
for _, tc := range []struct{ name, value string }{
|
wantStartStopped(t, []struct{ name, value string }{
|
||||||
{listenAddr, "8080"},
|
{listenAddr, "8080"}, {listenAddr, ":http"}, {listenAddr, ":65536"},
|
||||||
{listenAddr, ":http"},
|
|
||||||
{listenAddr, ":65536"},
|
|
||||||
{upstreamURL, "127.0.0.1:8081"},
|
{upstreamURL, "127.0.0.1:8081"},
|
||||||
{upstreamURL, "ftp://127.0.0.1:8081"},
|
{upstreamURL, "ftp://127.0.0.1:8081"},
|
||||||
{upstreamURL, "http://"},
|
{upstreamURL, "http://"},
|
||||||
@@ -213,6 +583,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{upstreamURL, "http://127.0.0.1:8081/app"},
|
{upstreamURL, "http://127.0.0.1:8081/app"},
|
||||||
{upstreamURL, "http://127.0.0.1:8081/?a=1"},
|
{upstreamURL, "http://127.0.0.1:8081/?a=1"},
|
||||||
{upstreamURL, "http://user:secret@127.0.0.1:8081"},
|
{upstreamURL, "http://user:secret@127.0.0.1:8081"},
|
||||||
|
{mode, "Observe"}, {mode, "block"}, {mode, ""},
|
||||||
{trustedProxies, "10.0.0.0/33"},
|
{trustedProxies, "10.0.0.0/33"},
|
||||||
{trustedProxies, "traefik"},
|
{trustedProxies, "traefik"},
|
||||||
{trustedProxies, "10.0.0.0/8,,192.168.0.0/16"},
|
{trustedProxies, "10.0.0.0/8,,192.168.0.0/16"},
|
||||||
@@ -220,8 +591,9 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{allowNets, "192.0.2.0/24,monitoring"},
|
{allowNets, "192.0.2.0/24,monitoring"},
|
||||||
{rateLimitExemptNets, "2001:db8::/129"},
|
{rateLimitExemptNets, "2001:db8::/129"},
|
||||||
{denyNets, "198.51.100.0/24,"},
|
{denyNets, "198.51.100.0/24,"},
|
||||||
{clientRequestTimeout, "60"},
|
{clientRequestTimeout, "60"}, {clientRequestTimeout, ""},
|
||||||
{clientRequestTimeout, ""},
|
{clientIdleTimeout, "0s"},
|
||||||
|
{clientIdleTimeout, "2 minutes"},
|
||||||
{clientResponseTimeout, "1y"},
|
{clientResponseTimeout, "1y"},
|
||||||
{upstreamRequestTimeout, "-1s"},
|
{upstreamRequestTimeout, "-1s"},
|
||||||
{upstreamResponseTimeout, "0s"},
|
{upstreamResponseTimeout, "0s"},
|
||||||
@@ -236,8 +608,8 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{rateLimitPerMinute, "1K"},
|
{rateLimitPerMinute, "1K"},
|
||||||
{rateLimitPerHour, "0"},
|
{rateLimitPerHour, "0"},
|
||||||
{rateLimitPerHour, "1.5"},
|
{rateLimitPerHour, "1.5"},
|
||||||
{rateLimitPerDay, "-1"},
|
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
|
||||||
{rateLimitPerDay, "lots"},
|
{rateLimitExemptPaths, "/assets/,,/static/"},
|
||||||
{deniedCountries, "nk"},
|
{deniedCountries, "nk"},
|
||||||
{deniedCountries, "kp,,ir"},
|
{deniedCountries, "kp,,ir"},
|
||||||
{deniedCountries, "prk"},
|
{deniedCountries, "prk"},
|
||||||
@@ -250,7 +622,37 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
{allowedCountries, "uk"},
|
{allowedCountries, "uk"},
|
||||||
{allowedCountries, "zz"},
|
{allowedCountries, "zz"},
|
||||||
{allowedCountries, "de,germany"},
|
{allowedCountries, "de,germany"},
|
||||||
} {
|
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
|
||||||
|
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
|
||||||
|
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
|
||||||
|
{logRequestHeaders, "host"}, {logRequestHeaders, "accept,Host"},
|
||||||
|
{logRequestHeaders, "transfer-encoding"}, {logRequestHeaders, "TRANSFER-ENCODING"},
|
||||||
|
{rulesEnabled, "yes"}, {rulesEnabled, "True"},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInvalidBanOrStateValueStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
wantStartStopped(t, []struct{ name, value string }{
|
||||||
|
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
|
||||||
|
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
|
||||||
|
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
|
||||||
|
{maxBanDuration, off}, {maxBanDuration, "1w"}, {attackBanDuration, off},
|
||||||
|
{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"},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantStartStopped checks that each setting, set to its value, stops the
|
||||||
|
// start with an error that names the setting.
|
||||||
|
func wantStartStopped(t *testing.T, invalid []struct{ name, value string }) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for _, tc := range invalid {
|
||||||
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -266,6 +668,187 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHostOrTransferEncodingStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Only Host's message points to the field host.
|
||||||
|
for value, want := range map[string]string{
|
||||||
|
"Host": `"Host" is taken out of every request by Go's HTTP server, ` +
|
||||||
|
"so it can never be logged; the request's host is the field host",
|
||||||
|
"transfer-encoding": `"transfer-encoding" is taken out of every ` +
|
||||||
|
"request by Go's HTTP server, so it can never be logged",
|
||||||
|
} {
|
||||||
|
t.Run(value, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{logRequestHeaders: value}.lookupEnv)
|
||||||
|
if err == nil || err.Error() != logRequestHeaders+": "+want {
|
||||||
|
t.Errorf("error %v, want %s: %s", err, logRequestHeaders, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Characters are counted, not bytes: each é takes two.
|
||||||
|
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
|
||||||
|
t.Run(value, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv)
|
||||||
|
|
||||||
|
want := metricsToken + ": is shorter than 32 characters"
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenIsLoggedMasked(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{metricsToken: token})
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
if strings.Contains(out.String(), token) ||
|
||||||
|
!strings.Contains(out.String(), `"`+metricsToken+`":"********"`) {
|
||||||
|
t.Errorf("the token is not logged masked: %s", out.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSettingFromFileLosesOneNewlineAndNoMore(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for contents, want := range map[string]string{
|
||||||
|
token: token,
|
||||||
|
token + "\n": token,
|
||||||
|
token + "\n\n": token + "\n",
|
||||||
|
token + " \n": token + " ",
|
||||||
|
} {
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
metricsToken + "_FILE": writeFile(t, contents),
|
||||||
|
})
|
||||||
|
if cfg.MetricsToken != want {
|
||||||
|
t.Errorf("file holding %q gave %s %q, want %q", contents, metricsToken,
|
||||||
|
cfg.MetricsToken, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSettingFromFileIsCheckedAsTheSettingItself(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
requestMaxBytes + "_FILE": writeFile(t, "lots\n"),
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want := requestMaxBytes + `: "lots" is not a size such as 512K, 100M or 5G, or off`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = config.FromEnvironment(environment{
|
||||||
|
logRemoteURL: remoteURL,
|
||||||
|
instanceName: instance,
|
||||||
|
logRemoteAppName + "_FILE": writeFile(t, "my app\n"),
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want = logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
|
||||||
|
`characters without a space, such as gitea`
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSettingAndItsFileBothSetStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
_, err := config.FromEnvironment(environment{
|
||||||
|
metricsToken: token,
|
||||||
|
metricsToken + "_FILE": writeFile(t, token),
|
||||||
|
}.lookupEnv)
|
||||||
|
|
||||||
|
want := metricsToken + ": is set, and so is " + metricsToken +
|
||||||
|
"_FILE; set only one of them"
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnreadableSettingFileStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
for _, path := range []string{filepath.Join(dir, "missing"), dir} {
|
||||||
|
_, err := config.FromEnvironment(environment{metricsToken + "_FILE": path}.lookupEnv)
|
||||||
|
|
||||||
|
want := metricsToken + "_FILE: cannot be read: "
|
||||||
|
if err == nil || !strings.HasPrefix(err.Error(), want) {
|
||||||
|
t.Errorf("%s: error %v, want one starting %s", path, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTokenFromFileIsLoggedMaskedWithTheFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
path := writeFile(t, token+"\n")
|
||||||
|
cfg := fromEnvironment(t, environment{metricsToken + "_FILE": path})
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
|
||||||
|
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||||
|
|
||||||
|
var line struct {
|
||||||
|
Settings map[string]string `json:"settings"`
|
||||||
|
}
|
||||||
|
|
||||||
|
err := json.Unmarshal(out.Bytes(), &line)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(out.String(), token) ||
|
||||||
|
line.Settings[metricsToken] != "********" ||
|
||||||
|
line.Settings[metricsToken+"_FILE"] != path {
|
||||||
|
t.Errorf("the token is not logged masked, with its file %s: %s", path,
|
||||||
|
out.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoteLogCAFileIsNotReadAsAFileInItsTurn(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
cfg := fromEnvironment(t, environment{
|
||||||
|
logRemoteTLSCAFile + "_FILE": writeFile(t, "/nonexistent/ca.pem\n"),
|
||||||
|
})
|
||||||
|
if cfg.LogRemoteTLSCAs != nil {
|
||||||
|
t.Errorf("%s_FILE gave certificates", logRemoteTLSCAFile)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeFile writes contents to a file in a directory of its own, removed
|
||||||
|
// when the test ends, and returns the file's path.
|
||||||
|
func writeFile(t *testing.T, contents string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "setting")
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte(contents), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
func TestLogsEachSettingWithItsValue(t *testing.T) {
|
func TestLogsEachSettingWithItsValue(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -284,11 +867,16 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
hostname, _ := os.Hostname()
|
||||||
|
|
||||||
want := map[string]string{
|
want := map[string]string{
|
||||||
listenAddr: ":8080",
|
listenAddr: ":8080",
|
||||||
upstreamURL: "http://127.0.0.1:8081",
|
upstreamURL: "http://127.0.0.1:8081",
|
||||||
|
mode: "enforce",
|
||||||
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
||||||
clientRequestTimeout: "45s",
|
clientRequestTimeout: "45s",
|
||||||
|
clientHeaderMaxBytes: "32K",
|
||||||
|
clientIdleTimeout: "120s",
|
||||||
clientResponseTimeout: "30m",
|
clientResponseTimeout: "30m",
|
||||||
upstreamRequestTimeout: "60s",
|
upstreamRequestTimeout: "60s",
|
||||||
upstreamResponseTimeout: "30m",
|
upstreamResponseTimeout: "30m",
|
||||||
@@ -300,8 +888,30 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
|||||||
rateLimitPerMinute: "1000",
|
rateLimitPerMinute: "1000",
|
||||||
rateLimitPerHour: "10000",
|
rateLimitPerHour: "10000",
|
||||||
rateLimitPerDay: "50000",
|
rateLimitPerDay: "50000",
|
||||||
|
rateLimitExemptPaths: "",
|
||||||
deniedCountries: "",
|
deniedCountries: "",
|
||||||
allowedCountries: "",
|
allowedCountries: "",
|
||||||
|
banResponse: "403",
|
||||||
|
limitBanDuration: "1h",
|
||||||
|
limitBanRepeatWindow: "24h",
|
||||||
|
maxBanDuration: "7d",
|
||||||
|
attackBanDuration: "7d",
|
||||||
|
maxBans: "5000",
|
||||||
|
banScopeV4Prefix: "32",
|
||||||
|
stateDir: "/var/lib/smallwebwaf",
|
||||||
|
stateWriteDelay: "10s",
|
||||||
|
stateCounterInterval: "15m",
|
||||||
|
metricsToken: "",
|
||||||
|
metricsTopN: "50",
|
||||||
|
instanceName: hostname,
|
||||||
|
logRequestHeaders: defaultLogRequestHeaders,
|
||||||
|
rulesDir: "/etc/smallwebwaf/rules.d",
|
||||||
|
rulesEnabled: "true",
|
||||||
|
logRemoteURL: "",
|
||||||
|
logRemoteTLSCAFile: "",
|
||||||
|
logRemoteBuffer: "10000",
|
||||||
|
logRemoteFacility: "local0",
|
||||||
|
logRemoteAppName: hostname,
|
||||||
}
|
}
|
||||||
if !maps.Equal(line.Settings, want) {
|
if !maps.Equal(line.Settings, want) {
|
||||||
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
|
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
|
||||||
@@ -313,7 +923,10 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
if got.ListenAddr != want.ListenAddr ||
|
if got.ListenAddr != want.ListenAddr ||
|
||||||
|
got.Observe != want.Observe ||
|
||||||
got.ClientRequestTimeout != want.ClientRequestTimeout ||
|
got.ClientRequestTimeout != want.ClientRequestTimeout ||
|
||||||
|
got.ClientRequestHeaderMaxBytes != want.ClientRequestHeaderMaxBytes ||
|
||||||
|
got.ClientIdleTimeout != want.ClientIdleTimeout ||
|
||||||
got.ClientResponseTimeout != want.ClientResponseTimeout ||
|
got.ClientResponseTimeout != want.ClientResponseTimeout ||
|
||||||
got.UpstreamRequestTimeout != want.UpstreamRequestTimeout ||
|
got.UpstreamRequestTimeout != want.UpstreamRequestTimeout ||
|
||||||
got.UpstreamResponseTimeout != want.UpstreamResponseTimeout ||
|
got.UpstreamResponseTimeout != want.UpstreamResponseTimeout ||
|
||||||
@@ -324,6 +937,40 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
|||||||
got.RateLimitPerDay != want.RateLimitPerDay {
|
got.RateLimitPerDay != want.RateLimitPerDay {
|
||||||
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
|
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
wantBanSettings(t, got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantBanSettings checks the settings for bans, the state files, the
|
||||||
|
// metrics and the rule 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.AttackBanDuration != want.AttackBanDuration ||
|
||||||
|
got.MaxBans != want.MaxBans ||
|
||||||
|
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
|
||||||
|
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.RulesDir != want.RulesDir || got.RulesEnabled != want.RulesEnabled {
|
||||||
|
t.Errorf("rule files in %q, on: %t, want %q, %t",
|
||||||
|
got.RulesDir, got.RulesEnabled, want.RulesDir, want.RulesEnabled)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.StateDir != want.StateDir ||
|
||||||
|
got.StateWriteDelay != want.StateWriteDelay ||
|
||||||
|
got.StateCounterInterval != want.StateCounterInterval {
|
||||||
|
t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.MetricsToken != want.MetricsToken || got.MetricsTopN != want.MetricsTopN {
|
||||||
|
t.Errorf("metrics token %q and top %d, want %q and %d",
|
||||||
|
got.MetricsToken, got.MetricsTopN, want.MetricsToken, want.MetricsTopN)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantNetblocks checks a list of netblocks.
|
// wantNetblocks checks a list of netblocks.
|
||||||
|
|||||||
@@ -0,0 +1,9 @@
|
|||||||
|
package lookup
|
||||||
|
|
||||||
|
import "net/http"
|
||||||
|
|
||||||
|
// SetTransport has g's requests to GeoJS go through transport instead of
|
||||||
|
// the network.
|
||||||
|
func (g *GeoJS) SetTransport(transport http.RoundTripper) {
|
||||||
|
g.httpClient.Transport = transport
|
||||||
|
}
|
||||||
+84
-13
@@ -1,6 +1,7 @@
|
|||||||
// Package lookup looks up each client's country through the GeoJS web
|
// Package lookup looks up each client's country through the GeoJS web
|
||||||
// service, and keeps the answers in memory, for at most 100,000 clients
|
// service, and keeps the answers in memory, for at most 100,000 clients
|
||||||
// and for 7 days each.
|
// and for 7 days each. The answers are written to lookups.json and read
|
||||||
|
// from it by the state package.
|
||||||
package lookup
|
package lookup
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -12,11 +13,13 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
// URL is GeoJS's country endpoint. Asked about several addresses at once,
|
// URL is GeoJS's country endpoint. Asked about several addresses at once,
|
||||||
@@ -62,6 +65,9 @@ type Params struct {
|
|||||||
Now func() time.Time
|
Now func() time.Time
|
||||||
// ProcessLog receives GeoJS's failures.
|
// ProcessLog receives GeoJS's failures.
|
||||||
ProcessLog *slog.Logger
|
ProcessLog *slog.Logger
|
||||||
|
// Metrics count the requests to GeoJS, those that failed, and the
|
||||||
|
// clients that go without an answer.
|
||||||
|
Metrics *metrics.Metrics
|
||||||
}
|
}
|
||||||
|
|
||||||
// GeoJS looks up clients' countries through GeoJS. At most one request
|
// GeoJS looks up clients' countries through GeoJS. At most one request
|
||||||
@@ -71,12 +77,13 @@ type GeoJS struct {
|
|||||||
url string
|
url string
|
||||||
now func() time.Time
|
now func() time.Time
|
||||||
processLog *slog.Logger
|
processLog *slog.Logger
|
||||||
|
metrics *metrics.Metrics
|
||||||
// httpClient follows no redirect, so that visitors' addresses go to
|
// httpClient follows no redirect, so that visitors' addresses go to
|
||||||
// GeoJS alone: a redirect is a failure.
|
// GeoJS alone: a redirect is a failure.
|
||||||
httpClient *http.Client
|
httpClient *http.Client
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
answers *simplelru.LRU[netip.Prefix, answer]
|
answers *simplelru.LRU[netip.Prefix, *Answer]
|
||||||
// waiting are the clients without an answer: those to ask GeoJS about,
|
// waiting are the clients without an answer: those to ask GeoJS about,
|
||||||
// and those it is being asked about.
|
// and those it is being asked about.
|
||||||
waiting map[netip.Prefix]*wait
|
waiting map[netip.Prefix]*wait
|
||||||
@@ -88,11 +95,14 @@ type GeoJS struct {
|
|||||||
retryAt time.Time
|
retryAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
// answer is what GeoJS said about a client: its country, "" when GeoJS
|
// Answer is what GeoJS said about a client, as lookups.json holds it: its
|
||||||
// cannot place it, and when GeoJS said so.
|
// country, "" when GeoJS cannot place it, when GeoJS said so, and when
|
||||||
type answer struct {
|
// the answer was last used.
|
||||||
country string
|
type Answer struct {
|
||||||
received time.Time
|
Client netip.Prefix `json:"client"`
|
||||||
|
Country string `json:"country"`
|
||||||
|
Answered time.Time `json:"answered"`
|
||||||
|
Used time.Time `json:"used"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// wait is a client waiting for its answer.
|
// wait is a client waiting for its answer.
|
||||||
@@ -107,7 +117,7 @@ type wait struct {
|
|||||||
|
|
||||||
// New returns a GeoJS with no answer kept yet.
|
// New returns a GeoJS with no answer kept yet.
|
||||||
func New(params Params) *GeoJS {
|
func New(params Params) *GeoJS {
|
||||||
answers, err := simplelru.NewLRU[netip.Prefix, answer](maxAnswers, nil)
|
answers, err := simplelru.NewLRU[netip.Prefix, *Answer](maxAnswers, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
@@ -116,6 +126,7 @@ func New(params Params) *GeoJS {
|
|||||||
url: params.URL,
|
url: params.URL,
|
||||||
now: params.Now,
|
now: params.Now,
|
||||||
processLog: params.ProcessLog,
|
processLog: params.ProcessLog,
|
||||||
|
metrics: params.Metrics,
|
||||||
httpClient: &http.Client{
|
httpClient: &http.Client{
|
||||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||||
return http.ErrUseLastResponse
|
return http.ErrUseLastResponse
|
||||||
@@ -155,6 +166,9 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
|||||||
defer g.mu.Unlock()
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
country, found := g.kept(client)
|
country, found := g.kept(client)
|
||||||
|
if !found {
|
||||||
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
|
}
|
||||||
|
|
||||||
w, waiting := g.waiting[client]
|
w, waiting := g.waiting[client]
|
||||||
if !found && waiting {
|
if !found && waiting {
|
||||||
@@ -164,6 +178,49 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
|
|||||||
return country
|
return country
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Snapshot returns every answer kept, sorted by client, as lookups.json
|
||||||
|
// lists them.
|
||||||
|
func (g *GeoJS) Snapshot() []Answer {
|
||||||
|
g.mu.Lock()
|
||||||
|
|
||||||
|
answers := make([]Answer, 0, g.answers.Len())
|
||||||
|
for _, kept := range g.answers.Values() {
|
||||||
|
answers = append(answers, *kept)
|
||||||
|
}
|
||||||
|
|
||||||
|
g.mu.Unlock()
|
||||||
|
|
||||||
|
slices.SortFunc(answers, func(a, b Answer) int {
|
||||||
|
return a.Client.Compare(b.Client)
|
||||||
|
})
|
||||||
|
|
||||||
|
return answers
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load keeps answers read from lookups.json, in place of the answers it
|
||||||
|
// keeps, in the order they were last used, so that the one used longest
|
||||||
|
// ago is dropped first. Answers GeoJS gave keepFor ago or more are
|
||||||
|
// dropped.
|
||||||
|
func (g *GeoJS) Load(answers []Answer) {
|
||||||
|
answers = slices.Clone(answers)
|
||||||
|
slices.SortStableFunc(answers, func(a, b Answer) int {
|
||||||
|
return a.Used.Compare(b.Used)
|
||||||
|
})
|
||||||
|
|
||||||
|
g.mu.Lock()
|
||||||
|
defer g.mu.Unlock()
|
||||||
|
|
||||||
|
g.answers.Purge()
|
||||||
|
|
||||||
|
now := g.now()
|
||||||
|
|
||||||
|
for _, answer := range answers {
|
||||||
|
if now.Sub(answer.Answered) < keepFor {
|
||||||
|
g.answers.Add(answer.Client, &answer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// answerOrWait returns client's kept answer if it has one. Otherwise it
|
// answerOrWait returns client's kept answer if it has one. Otherwise it
|
||||||
// puts the client among those waiting if there is room, has GeoJS asked
|
// puts the client among those waiting if there is room, has GeoJS asked
|
||||||
// about them if it can be, and returns what to wait on for the answer, or
|
// about them if it can be, and returns what to wait on for the answer, or
|
||||||
@@ -188,6 +245,8 @@ func (g *GeoJS) answerOrWait(
|
|||||||
g.ask(ctx)
|
g.ask(ctx)
|
||||||
|
|
||||||
if w == nil {
|
if w == nil {
|
||||||
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
|
|
||||||
return "", nil // too many clients wait already
|
return "", nil // too many clients wait already
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -197,21 +256,27 @@ func (g *GeoJS) answerOrWait(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if w.late {
|
if w.late {
|
||||||
|
g.metrics.GeoJSUnanswered.Inc()
|
||||||
|
|
||||||
return "", nil
|
return "", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return "", w.asked
|
return "", w.asked
|
||||||
}
|
}
|
||||||
|
|
||||||
// kept returns client's answer, if one was received less than keepFor
|
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
|
||||||
// ago.
|
// and notes that it was used.
|
||||||
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
|
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
|
||||||
|
now := g.now()
|
||||||
|
|
||||||
kept, found := g.answers.Get(client)
|
kept, found := g.answers.Get(client)
|
||||||
if !found || g.now().Sub(kept.received) >= keepFor {
|
if !found || now.Sub(kept.Answered) >= keepFor {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
|
|
||||||
return kept.country, true
|
kept.Used = now
|
||||||
|
|
||||||
|
return kept.Country, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// ask starts asking GeoJS about the waiting clients, unless a request to
|
// ask starts asking GeoJS about the waiting clients, unless a request to
|
||||||
@@ -293,7 +358,9 @@ func (g *GeoJS) keep(
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
g.answers.Add(client, answer{country: country, received: now})
|
g.answers.Add(client, &Answer{
|
||||||
|
Client: client, Country: country, Answered: now, Used: now,
|
||||||
|
})
|
||||||
close(g.waiting[client].asked)
|
close(g.waiting[client].asked)
|
||||||
delete(g.waiting, client)
|
delete(g.waiting, client)
|
||||||
}
|
}
|
||||||
@@ -303,6 +370,8 @@ func (g *GeoJS) keep(
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
g.metrics.GeoJSFailures.Inc()
|
||||||
|
|
||||||
g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay),
|
g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay),
|
||||||
maxRetryDelay)
|
maxRetryDelay)
|
||||||
g.retryAt = now.Add(g.retryDelay)
|
g.retryAt = now.Add(g.retryDelay)
|
||||||
@@ -347,6 +416,8 @@ func (g *GeoJS) request(
|
|||||||
|
|
||||||
req.URL.RawQuery = "ip=" + strings.Join(addrs, ",")
|
req.URL.RawQuery = "ip=" + strings.Join(addrs, ",")
|
||||||
|
|
||||||
|
g.metrics.GeoJSRequests.Inc()
|
||||||
|
|
||||||
res, err := g.httpClient.Do(req)
|
res, err := g.httpClient.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Do's error names the URL, and so the visitors' addresses, which
|
// Do's error names the URL, and so the visitors' addresses, which
|
||||||
|
|||||||
+303
-225
@@ -10,9 +10,12 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus/testutil"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -26,80 +29,87 @@ const (
|
|||||||
leftOut = "203.0.113.7"
|
leftOut = "203.0.113.7"
|
||||||
// timeout is how long a new client waits for its answer.
|
// timeout is how long a new client waits for its answer.
|
||||||
timeout = time.Second
|
timeout = time.Second
|
||||||
// waitLimit bounds how long a test waits for what should happen.
|
|
||||||
waitLimit = 10 * time.Second
|
|
||||||
// pollInterval is how often a test looks again.
|
|
||||||
pollInterval = 10 * time.Millisecond
|
|
||||||
// week is how long an answer is kept.
|
// week is how long an answer is kept.
|
||||||
week = 7 * 24 * time.Hour
|
week = 7 * 24 * time.Hour
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// The tests that have GeoJS asked run in a synctest bubble, where the time
|
||||||
|
// package runs on a clock of the test's own: a wait lasts exactly as long
|
||||||
|
// as it should, however slowly the test process runs, and synctest.Wait
|
||||||
|
// returns once g has done all it can before time passes. The stand-in for
|
||||||
|
// GeoJS answers without the network, since a request waiting on the
|
||||||
|
// network would keep that clock from moving on.
|
||||||
|
|
||||||
func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) {
|
func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
placed := netip.MustParsePrefix("203.0.113.9/32")
|
geojs, clock, g := start()
|
||||||
notPlaced := netip.MustParsePrefix(unplaced + "/32")
|
placed := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
notPlaced := netip.MustParsePrefix(unplaced + "/32")
|
||||||
|
|
||||||
wantCountry(t, g, placed, germany)
|
wantCountry(t, g, placed, germany)
|
||||||
wantCountry(t, g, notPlaced, "")
|
wantCountry(t, g, notPlaced, "")
|
||||||
wantRequests(t, geojs, 2)
|
wantRequests(t, geojs, 2)
|
||||||
|
|
||||||
// An answer without a country is kept too.
|
// An answer without a country is kept too.
|
||||||
clock.advance(week - time.Second)
|
clock.advance(week - time.Second)
|
||||||
wantCountry(t, g, placed, germany)
|
wantCountry(t, g, placed, germany)
|
||||||
wantCountry(t, g, notPlaced, "")
|
wantCountry(t, g, notPlaced, "")
|
||||||
wantRequests(t, geojs, 2)
|
wantRequests(t, geojs, 2)
|
||||||
|
|
||||||
clock.advance(time.Second)
|
clock.advance(time.Second)
|
||||||
wantCountry(t, g, placed, germany)
|
wantCountry(t, g, placed, germany)
|
||||||
wantRequests(t, geojs, 3)
|
wantRequests(t, geojs, 3)
|
||||||
wantAsked(t, geojs, 2, "203.0.113.9")
|
wantAsked(t, geojs, 2, "203.0.113.9")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
|
func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
geojs, clock, g := start()
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
|
||||||
// The client comes while GeoJS is asked about an earlier client, which
|
// The client comes while GeoJS is asked about an earlier client, which
|
||||||
// it answers most of a second later. It is then asked about the client
|
// it answers most of a second later. It is then asked about the client
|
||||||
// and does not answer: that request is abandoned a second after it
|
// and does not answer: that request is abandoned a second after it
|
||||||
// began, well after the client's wait is over.
|
// began, well after the client's wait is over.
|
||||||
geojs.set(answeringSlowly)
|
geojs.set(answeringSlowly)
|
||||||
|
|
||||||
var earlier sync.WaitGroup
|
var earlier sync.WaitGroup
|
||||||
|
|
||||||
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
|
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
|
||||||
defer earlier.Wait()
|
defer earlier.Wait()
|
||||||
|
|
||||||
waitForRequests(t, geojs, 1)
|
waitForRequests(t, geojs, 1)
|
||||||
geojs.set(hanging)
|
geojs.set(hanging)
|
||||||
|
|
||||||
began := time.Now()
|
began := time.Now()
|
||||||
|
|
||||||
wantCountry(t, g, client, "")
|
wantCountry(t, g, client, "")
|
||||||
|
|
||||||
took := time.Since(began)
|
took := time.Since(began)
|
||||||
if took < timeout || took > timeout+timeout/2 {
|
if took != timeout {
|
||||||
t.Errorf("waited %s for the answer, want %s", took, timeout)
|
t.Errorf("waited %s for the answer, want %s", took, timeout)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Its next request does not wait.
|
// Its next request does not wait.
|
||||||
began = time.Now()
|
began = time.Now()
|
||||||
|
|
||||||
wantCountry(t, g, client, "")
|
wantCountry(t, g, client, "")
|
||||||
|
|
||||||
took = time.Since(began)
|
took = time.Since(began)
|
||||||
if took > timeout/2 {
|
if took != 0 {
|
||||||
t.Errorf("waited %s again, want no wait", took)
|
t.Errorf("waited %s again, want no wait", took)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Once GeoJS answers, the client is asked about again in the
|
// Once GeoJS answers, the client is asked about again in the
|
||||||
// background, and has its country.
|
// background, and has its country.
|
||||||
geojs.set(answering)
|
geojs.set(answering)
|
||||||
waitForCountry(t, g, clock, client, germany)
|
waitForCountry(t, g, clock, client, germany)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
||||||
@@ -118,33 +128,35 @@ func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
|||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
other := netip.MustParsePrefix("203.0.113.1/32")
|
geojs, clock, g := start()
|
||||||
client := netip.MustParsePrefix(leftOut + "/32")
|
other := netip.MustParsePrefix("203.0.113.1/32")
|
||||||
|
client := netip.MustParsePrefix(leftOut + "/32")
|
||||||
|
|
||||||
// GeoJS fails, and is left alone for a second while the client
|
// GeoJS fails, and is left alone for a second while the client
|
||||||
// comes too, so that the next request asks about both.
|
// comes too, so that the next request asks about both.
|
||||||
geojs.set(failing)
|
geojs.set(failing)
|
||||||
wantCountry(t, g, other, "")
|
wantCountry(t, g, other, "")
|
||||||
wantCountry(t, g, client, "")
|
wantCountry(t, g, client, "")
|
||||||
|
|
||||||
geojs.set(tc.answers)
|
geojs.set(tc.answers)
|
||||||
clock.advance(time.Second)
|
clock.advance(time.Second)
|
||||||
wantCountry(t, g, other, "")
|
wantCountry(t, g, other, "")
|
||||||
waitForRequests(t, geojs, 2)
|
waitForRequests(t, geojs, 2)
|
||||||
|
|
||||||
// The answer counts as a failure, and the client is asked about
|
// The answer counts as a failure, and the client is asked about
|
||||||
// again, with the other client only if the answer left it out too.
|
// again, with the other client only if the answer left it out too.
|
||||||
geojs.set(answering)
|
geojs.set(answering)
|
||||||
waitForCountry(t, g, clock, client, germany)
|
waitForCountry(t, g, clock, client, germany)
|
||||||
wantCountry(t, g, other, germany)
|
wantCountry(t, g, other, germany)
|
||||||
wantRequests(t, geojs, 3)
|
wantRequests(t, geojs, 3)
|
||||||
|
|
||||||
if tc.named {
|
if tc.named {
|
||||||
wantAsked(t, geojs, 2, leftOut)
|
wantAsked(t, geojs, 2, leftOut)
|
||||||
} else {
|
} else {
|
||||||
wantAsked(t, geojs, 2, leftOut, "203.0.113.1")
|
wantAsked(t, geojs, 2, leftOut, "203.0.113.1")
|
||||||
}
|
}
|
||||||
|
})
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -152,173 +164,225 @@ func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
|||||||
func TestRedirectCountsAsFailure(t *testing.T) {
|
func TestRedirectCountsAsFailure(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, _, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
geojs.set(redirecting)
|
geojs, _, g := start()
|
||||||
|
geojs.set(redirecting)
|
||||||
|
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
||||||
wantRequests(t, geojs, 1)
|
wantRequests(t, geojs, 1)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCountryIsKeptInCapitals(t *testing.T) {
|
func TestCountryIsKeptInCapitals(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, _, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
geojs.set(answeringInLowerCase)
|
geojs, _, g := start()
|
||||||
|
geojs.set(answeringInLowerCase)
|
||||||
|
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany)
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
var log strings.Builder
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
var log strings.Builder
|
||||||
|
|
||||||
// Nothing listens on port 1, so asking GeoJS fails.
|
// GeoJS does not answer, so the request to it is abandoned, and fails.
|
||||||
g := lookup.New(lookup.Params{
|
geojs := &standIn{answers: hanging}
|
||||||
URL: "http://127.0.0.1:1",
|
g := lookup.New(lookup.Params{
|
||||||
Now: time.Now,
|
URL: lookup.URL,
|
||||||
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
|
Now: time.Now,
|
||||||
|
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
|
||||||
|
Metrics: metrics.New(1),
|
||||||
|
})
|
||||||
|
g.SetTransport(geojs)
|
||||||
|
|
||||||
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
logged := log.String()
|
||||||
|
if !strings.Contains(logged, "asking GeoJS failed") ||
|
||||||
|
strings.Contains(logged, "203.0.113.9") {
|
||||||
|
t.Errorf("logged %q, want the failure without the address asked about", logged)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
|
||||||
|
|
||||||
logged := log.String()
|
|
||||||
if !strings.Contains(logged, "asking GeoJS failed") ||
|
|
||||||
strings.Contains(logged, "203.0.113.9") {
|
|
||||||
t.Errorf("logged %q, want the failure without the address asked about", logged)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
|
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
geojs, clock, g := start()
|
||||||
|
|
||||||
// GeoJS fails, and is then left alone for a second, while three more
|
// GeoJS fails, and is then left alone for a second, while three more
|
||||||
// clients come. An IPv6 client is a /64, and GeoJS is asked about its
|
// clients come. An IPv6 client is a /64, and GeoJS is asked about its
|
||||||
// first address.
|
// first address.
|
||||||
geojs.set(failing)
|
geojs.set(failing)
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "")
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "")
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "")
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "")
|
||||||
wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "")
|
wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "")
|
||||||
wantRequests(t, geojs, 1)
|
wantRequests(t, geojs, 1)
|
||||||
|
|
||||||
geojs.set(answering)
|
geojs.set(answering)
|
||||||
clock.advance(time.Second)
|
clock.advance(time.Second)
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany)
|
wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany)
|
||||||
wantRequests(t, geojs, 2)
|
wantRequests(t, geojs, 2)
|
||||||
wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3")
|
wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) {
|
func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
clients := newClients()
|
geojs, clock, g := start()
|
||||||
kept := clients()
|
clients := newClients()
|
||||||
|
kept := clients()
|
||||||
|
|
||||||
wantCountry(t, g, kept, germany)
|
|
||||||
|
|
||||||
geojs.set(failing)
|
|
||||||
wantCountry(t, g, kept, germany)
|
|
||||||
wantRequests(t, geojs, 1)
|
|
||||||
|
|
||||||
// Each failure leaves GeoJS alone twice as long as the one before, up
|
|
||||||
// to five minutes. New clients meanwhile count as not found, and the
|
|
||||||
// client with a kept answer still gets its country, without GeoJS being
|
|
||||||
// asked.
|
|
||||||
requests := 1
|
|
||||||
|
|
||||||
for _, delay := range []time.Duration{
|
|
||||||
time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second,
|
|
||||||
16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second,
|
|
||||||
256 * time.Second, 5 * time.Minute, 5 * time.Minute,
|
|
||||||
} {
|
|
||||||
wantCountry(t, g, clients(), "")
|
|
||||||
|
|
||||||
requests++
|
|
||||||
wantRequests(t, geojs, requests)
|
|
||||||
|
|
||||||
clock.advance(delay - time.Millisecond)
|
|
||||||
wantCountry(t, g, clients(), "")
|
|
||||||
wantCountry(t, g, kept, germany)
|
wantCountry(t, g, kept, germany)
|
||||||
wantRequests(t, geojs, requests)
|
|
||||||
|
|
||||||
clock.advance(time.Millisecond)
|
geojs.set(failing)
|
||||||
}
|
wantCountry(t, g, kept, germany)
|
||||||
|
wantRequests(t, geojs, 1)
|
||||||
|
|
||||||
// Once GeoJS answers again, it is asked about every client waiting.
|
// Each failure leaves GeoJS alone twice as long as the one before, up
|
||||||
geojs.set(answering)
|
// to five minutes. New clients meanwhile count as not found, and the
|
||||||
wantCountry(t, g, clients(), germany)
|
// client with a kept answer still gets its country, without GeoJS being
|
||||||
wantRequests(t, geojs, requests+1)
|
// asked.
|
||||||
|
requests := 1
|
||||||
|
|
||||||
asked := waitForRequests(t, geojs, requests+1)
|
for _, delay := range []time.Duration{
|
||||||
if len(asked[requests]) != 23 {
|
time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second,
|
||||||
t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests]))
|
16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second,
|
||||||
}
|
256 * time.Second, 5 * time.Minute, 5 * time.Minute,
|
||||||
|
} {
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
|
||||||
|
requests++
|
||||||
|
wantRequests(t, geojs, requests)
|
||||||
|
|
||||||
|
clock.advance(delay - time.Millisecond)
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
wantCountry(t, g, kept, germany)
|
||||||
|
wantRequests(t, geojs, requests)
|
||||||
|
|
||||||
|
clock.advance(time.Millisecond)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Once GeoJS answers again, it is asked about every client waiting.
|
||||||
|
geojs.set(answering)
|
||||||
|
wantCountry(t, g, clients(), germany)
|
||||||
|
wantRequests(t, geojs, requests+1)
|
||||||
|
|
||||||
|
asked := waitForRequests(t, geojs, requests+1)
|
||||||
|
if len(asked[requests]) != 23 {
|
||||||
|
t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests]))
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAtMost200AddressesInOneRequest(t *testing.T) {
|
func TestAtMost200AddressesInOneRequest(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
clients := newClients()
|
geojs, clock, g := start()
|
||||||
first := clients()
|
clients := newClients()
|
||||||
|
first := clients()
|
||||||
|
|
||||||
// 201 clients wait while GeoJS is left alone after a failure.
|
// 201 clients wait while GeoJS is left alone after a failure.
|
||||||
geojs.set(failing)
|
geojs.set(failing)
|
||||||
wantCountry(t, g, first, "")
|
wantCountry(t, g, first, "")
|
||||||
|
|
||||||
for range 200 {
|
for range 200 {
|
||||||
wantCountry(t, g, clients(), "")
|
wantCountry(t, g, clients(), "")
|
||||||
}
|
}
|
||||||
|
|
||||||
// The first one's next request has GeoJS asked again.
|
// The first one's next request has GeoJS asked again.
|
||||||
geojs.set(answering)
|
geojs.set(answering)
|
||||||
clock.advance(time.Second)
|
clock.advance(time.Second)
|
||||||
wantCountry(t, g, first, "")
|
wantCountry(t, g, first, "")
|
||||||
|
|
||||||
asked := waitForRequests(t, geojs, 3)
|
asked := waitForRequests(t, geojs, 3)
|
||||||
if len(asked[1]) != 200 || len(asked[2]) != 1 {
|
if len(asked[1]) != 200 || len(asked[2]) != 1 {
|
||||||
t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1",
|
t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1",
|
||||||
len(asked[1]), len(asked[2]))
|
len(asked[1]), len(asked[2]))
|
||||||
}
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAtMost10000ClientsWait(t *testing.T) {
|
func TestAtMost10000ClientsWait(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
geojs, clock, g := start(t)
|
synctest.Test(t, func(t *testing.T) {
|
||||||
clients := newClients()
|
geojs, clock, g := start()
|
||||||
first := clients()
|
clients := newClients()
|
||||||
|
first := clients()
|
||||||
|
|
||||||
// 10,000 clients wait while GeoJS is left alone after a failure, and
|
// 10,000 clients wait while GeoJS is left alone after a failure, and
|
||||||
// one more cannot join them.
|
// one more cannot join them.
|
||||||
geojs.set(failing)
|
geojs.set(failing)
|
||||||
wantCountry(t, g, first, "")
|
wantCountry(t, g, first, "")
|
||||||
|
|
||||||
for range 9999 {
|
for range 9999 {
|
||||||
wantCountry(t, g, clients(), "")
|
wantCountry(t, g, clients(), "")
|
||||||
}
|
|
||||||
|
|
||||||
extra := clients()
|
|
||||||
wantCountry(t, g, extra, "")
|
|
||||||
|
|
||||||
// The first one's next request has GeoJS asked about the 10,000, 200
|
|
||||||
// at a time, and not about the one more.
|
|
||||||
geojs.set(answering)
|
|
||||||
clock.advance(time.Second)
|
|
||||||
wantCountry(t, g, first, "")
|
|
||||||
|
|
||||||
asked := waitForRequests(t, geojs, 51)
|
|
||||||
for i, request := range asked {
|
|
||||||
if slices.Contains(request, extra.Addr().String()) {
|
|
||||||
t.Errorf("request %d asked about %s", i, extra.Addr())
|
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
// With room among those waiting, it is asked about.
|
extra := clients()
|
||||||
wantCountry(t, g, extra, germany)
|
wantCountry(t, g, extra, "")
|
||||||
|
|
||||||
|
// The first one's next request has GeoJS asked about the 10,000, 200
|
||||||
|
// at a time, and not about the one more.
|
||||||
|
geojs.set(answering)
|
||||||
|
clock.advance(time.Second)
|
||||||
|
wantCountry(t, g, first, "")
|
||||||
|
|
||||||
|
asked := waitForRequests(t, geojs, 51)
|
||||||
|
for i, request := range asked {
|
||||||
|
if slices.Contains(request, extra.Addr().String()) {
|
||||||
|
t.Errorf("request %d asked about %s", i, extra.Addr())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// With room among those waiting, it is asked about.
|
||||||
|
wantCountry(t, g, extra, germany)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
m := metrics.New(1)
|
||||||
|
g := lookup.New(lookup.Params{
|
||||||
|
URL: lookup.URL,
|
||||||
|
Now: time.Now,
|
||||||
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Metrics: m,
|
||||||
|
})
|
||||||
|
g.SetTransport(&standIn{answers: failing})
|
||||||
|
|
||||||
|
clients := newClients()
|
||||||
|
|
||||||
|
// GeoJS fails, so the first client goes without an answer, and GeoJS
|
||||||
|
// is left alone for a second, which does not pass in this test.
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
wantUnanswered(t, m, 1)
|
||||||
|
|
||||||
|
// Meanwhile each new client goes without one at once, while there is
|
||||||
|
// room for it among the 10,000 that may wait.
|
||||||
|
for range 9999 {
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
}
|
||||||
|
|
||||||
|
wantUnanswered(t, m, 10000)
|
||||||
|
|
||||||
|
// One more, for which there is no room, goes without one too.
|
||||||
|
wantCountry(t, g, clients(), "")
|
||||||
|
wantUnanswered(t, m, 10001)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// How the stand-in for GeoJS answers.
|
// How the stand-in for GeoJS answers.
|
||||||
@@ -337,13 +401,25 @@ const (
|
|||||||
// standIn is a stand-in for GeoJS. It notes the addresses each request
|
// standIn is a stand-in for GeoJS. It notes the addresses each request
|
||||||
// asks about.
|
// asks about.
|
||||||
type standIn struct {
|
type standIn struct {
|
||||||
server *httptest.Server
|
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
answers int
|
answers int
|
||||||
requests [][]string
|
requests [][]string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RoundTrip has the stand-in answer req, in place of the network. A request
|
||||||
|
// abandoned before the stand-in answers fails, as over the network.
|
||||||
|
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
answer := httptest.NewRecorder()
|
||||||
|
s.ServeHTTP(answer, req)
|
||||||
|
|
||||||
|
err := req.Context().Err()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return answer.Result(), nil
|
||||||
|
}
|
||||||
|
|
||||||
// ServeHTTP answers a request about the addresses in its ip parameter.
|
// ServeHTTP answers a request about the addresses in its ip parameter.
|
||||||
func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
addrs := strings.Split(r.URL.Query().Get("ip"), ",")
|
addrs := strings.Split(r.URL.Query().Get("ip"), ",")
|
||||||
@@ -422,7 +498,8 @@ func (s *standIn) asked() [][]string {
|
|||||||
return slices.Clone(s.requests)
|
return slices.Clone(s.requests)
|
||||||
}
|
}
|
||||||
|
|
||||||
// testClock is a clock the test sets.
|
// testClock is a clock the test sets. GeoJS tells the time by it, while
|
||||||
|
// waits run on the bubble's clock.
|
||||||
type testClock struct {
|
type testClock struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
now time.Time
|
now time.Time
|
||||||
@@ -444,21 +521,18 @@ func (c *testClock) advance(d time.Duration) {
|
|||||||
c.now = c.now.Add(d)
|
c.now = c.now.Add(d)
|
||||||
}
|
}
|
||||||
|
|
||||||
// start starts a stand-in for GeoJS that answers, and returns it, a
|
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
|
||||||
// clock, and a GeoJS asking it by that clock.
|
// asking the stand-in by that clock.
|
||||||
func start(t *testing.T) (*standIn, *testClock, *lookup.GeoJS) {
|
func start() (*standIn, *testClock, *lookup.GeoJS) {
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
geojs := &standIn{}
|
geojs := &standIn{}
|
||||||
geojs.server = httptest.NewServer(geojs)
|
|
||||||
t.Cleanup(geojs.server.Close)
|
|
||||||
|
|
||||||
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
|
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
|
||||||
g := lookup.New(lookup.Params{
|
g := lookup.New(lookup.Params{
|
||||||
URL: geojs.server.URL,
|
URL: lookup.URL,
|
||||||
Now: clock.Now,
|
Now: clock.Now,
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
Metrics: metrics.New(1),
|
||||||
})
|
})
|
||||||
|
g.SetTransport(geojs)
|
||||||
|
|
||||||
return geojs, clock, g
|
return geojs, clock, g
|
||||||
}
|
}
|
||||||
@@ -513,41 +587,45 @@ func wantAsked(t *testing.T, geojs *standIn, i int, want ...string) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// waitForRequests waits for GeoJS to have had count requests, and returns
|
// wantUnanswered checks how many requests m counts as having gone without
|
||||||
// the addresses each asked about.
|
// an answer from GeoJS.
|
||||||
|
func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := testutil.ToFloat64(m.GeoJSUnanswered)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("%v requests went without an answer, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitForRequests waits until g has done all it can before time passes,
|
||||||
|
// checks that GeoJS has had count requests, and returns the addresses each
|
||||||
|
// asked about.
|
||||||
func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string {
|
func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
deadline := time.Now().Add(waitLimit)
|
synctest.Wait()
|
||||||
for time.Now().Before(deadline) {
|
|
||||||
asked := geojs.asked()
|
|
||||||
if len(asked) >= count {
|
|
||||||
return asked
|
|
||||||
}
|
|
||||||
|
|
||||||
time.Sleep(pollInterval)
|
asked := geojs.asked()
|
||||||
|
if len(asked) != count {
|
||||||
|
t.Fatalf("GeoJS had %d requests, want %d", len(asked), count)
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Fatalf("fewer than %d requests to GeoJS after %s", count, waitLimit)
|
return asked
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// waitForCountry waits for g to give client the country want, moving the
|
// waitForCountry lets a request to GeoJS under way be abandoned, and moves
|
||||||
// clock on a minute at a time, so that GeoJS is asked again after a
|
// the clock on a minute, so that GeoJS may be asked again after a failure.
|
||||||
// failure.
|
// It then checks that client's next request does not wait but has it asked
|
||||||
|
// about again in the background, after which g gives it the country want.
|
||||||
func waitForCountry(
|
func waitForCountry(
|
||||||
t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string,
|
t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string,
|
||||||
) {
|
) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
deadline := time.Now().Add(waitLimit)
|
time.Sleep(timeout)
|
||||||
for g.Country(t.Context(), client) != want {
|
clock.advance(time.Minute)
|
||||||
if time.Now().After(deadline) {
|
wantCountry(t, g, client, "")
|
||||||
t.Fatalf("%s is not in %q after %s", client, want, waitLimit)
|
synctest.Wait()
|
||||||
}
|
wantCountry(t, g, client, want)
|
||||||
|
|
||||||
clock.advance(time.Minute)
|
|
||||||
time.Sleep(pollInterval)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,99 @@
|
|||||||
|
package lookup_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
_, clock, g := start()
|
||||||
|
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()
|
||||||
|
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()
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// other is the label under which the countries outside the busiest are
|
||||||
|
// counted.
|
||||||
|
const other = "other"
|
||||||
|
|
||||||
|
// countries are the metrics by the client's country, for requests whose
|
||||||
|
// client's country is known. The topN busiest countries, by their requests
|
||||||
|
// since the start, have series of their own, and the others are counted
|
||||||
|
// under other, so that there are never more than topN + 1 series. A
|
||||||
|
// country that drops out of the busiest loses its series, and its next
|
||||||
|
// requests are counted under other; one that becomes one of them gets a
|
||||||
|
// series that counts from then on. Each series therefore only ever goes
|
||||||
|
// up.
|
||||||
|
type countries struct {
|
||||||
|
topN int
|
||||||
|
|
||||||
|
requests *prometheus.CounterVec
|
||||||
|
requestBytes *prometheus.CounterVec
|
||||||
|
responseBytes *prometheus.CounterVec
|
||||||
|
// refused are the requests the country lists refused.
|
||||||
|
refused *prometheus.CounterVec
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
// seen is each country's requests since the start, by which the
|
||||||
|
// countries are ranked. GeoJS gives two-letter codes, so it holds at
|
||||||
|
// most a few hundred.
|
||||||
|
seen map[string]int64
|
||||||
|
// top are the countries with series of their own.
|
||||||
|
top map[string]bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// newCountries returns the metrics by country, with series of their own
|
||||||
|
// for the topN busiest countries.
|
||||||
|
func newCountries(topN int) *countries {
|
||||||
|
byCountry := []string{"country"}
|
||||||
|
|
||||||
|
return &countries{
|
||||||
|
topN: topN,
|
||||||
|
requests: counterVec("smallwebwaf_country_requests_total",
|
||||||
|
"Requests, by the client's country.", byCountry),
|
||||||
|
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
|
||||||
|
"Request body bytes, by the client's country.", byCountry),
|
||||||
|
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
|
||||||
|
"Response body bytes, by the client's country.", byCountry),
|
||||||
|
refused: counterVec("smallwebwaf_country_list_refusals_total",
|
||||||
|
"Requests the country lists refused, by the client's country.",
|
||||||
|
byCountry),
|
||||||
|
seen: map[string]int64{},
|
||||||
|
top: map[string]bool{},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// add counts a request from its log line, whose country is known.
|
||||||
|
func (c *countries) add(line *requestlog.Line) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
|
||||||
|
c.seen[line.Country]++
|
||||||
|
|
||||||
|
label := c.label(line.Country)
|
||||||
|
c.requests.WithLabelValues(label).Inc()
|
||||||
|
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
|
||||||
|
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
|
||||||
|
|
||||||
|
if line.Action == requestlog.ActionCountryDenied {
|
||||||
|
c.refused.WithLabelValues(label).Inc()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// label returns the label a request from country is counted under: the
|
||||||
|
// country while it is one of the busiest, other while it is not. A
|
||||||
|
// country busier than the least busy of them takes its place, and that
|
||||||
|
// country's series are dropped.
|
||||||
|
func (c *countries) label(country string) string {
|
||||||
|
if c.top[country] {
|
||||||
|
return country
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(c.top) < c.topN {
|
||||||
|
c.top[country] = true
|
||||||
|
|
||||||
|
return country
|
||||||
|
}
|
||||||
|
|
||||||
|
least := ""
|
||||||
|
|
||||||
|
for top := range c.top {
|
||||||
|
if least == "" || c.seen[top] < c.seen[least] {
|
||||||
|
least = top
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.seen[country] <= c.seen[least] {
|
||||||
|
return other
|
||||||
|
}
|
||||||
|
|
||||||
|
delete(c.top, least)
|
||||||
|
|
||||||
|
for _, vec := range []*prometheus.CounterVec{
|
||||||
|
c.requests, c.requestBytes, c.responseBytes, c.refused,
|
||||||
|
} {
|
||||||
|
vec.DeleteLabelValues(least)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.top[country] = true
|
||||||
|
|
||||||
|
return country
|
||||||
|
}
|
||||||
@@ -0,0 +1,335 @@
|
|||||||
|
// Package metrics keeps smallwebwaf's Prometheus metrics, as the "Metrics
|
||||||
|
// endpoint" section of SPEC.md lists them, and serves them in the
|
||||||
|
// Prometheus text format. No metric carries a client's address.
|
||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/collectors"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
|
||||||
|
type Metrics struct {
|
||||||
|
registry *prometheus.Registry
|
||||||
|
handler http.Handler
|
||||||
|
|
||||||
|
inFlight prometheus.Gauge
|
||||||
|
requests *prometheus.CounterVec
|
||||||
|
requestBytes *prometheus.CounterVec
|
||||||
|
responseBytes *prometheus.CounterVec
|
||||||
|
requestDuration prometheus.Histogram
|
||||||
|
upstreamDuration prometheus.Histogram
|
||||||
|
rateLimitHits *prometheus.CounterVec
|
||||||
|
sizeAndTimeLimitHits *prometheus.CounterVec
|
||||||
|
offences *prometheus.CounterVec
|
||||||
|
// ruleMatches are made by AddRules.
|
||||||
|
ruleMatches *prometheus.CounterVec
|
||||||
|
countries *countries
|
||||||
|
|
||||||
|
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
|
||||||
|
// that failed. GeoJSUnanswered are the requests whose client counted
|
||||||
|
// as coming from an unknown country because GeoJS had not answered
|
||||||
|
// about it in time.
|
||||||
|
GeoJSRequests prometheus.Counter
|
||||||
|
GeoJSFailures prometheus.Counter
|
||||||
|
GeoJSUnanswered prometheus.Counter
|
||||||
|
|
||||||
|
stateFileWrites *prometheus.CounterVec
|
||||||
|
stateFileWriteFailures *prometheus.CounterVec
|
||||||
|
stateFileLastWrite *prometheus.GaugeVec
|
||||||
|
stateFileSize *prometheus.GaugeVec
|
||||||
|
stateFileEditsTakenIn *prometheus.CounterVec
|
||||||
|
stateFileEditsSetAside *prometheus.CounterVec
|
||||||
|
}
|
||||||
|
|
||||||
|
// New returns the metrics, with the Go runtime's and the process's own.
|
||||||
|
// topN is how many countries get series of their own
|
||||||
|
// (SWWAF_METRICS_TOP_N).
|
||||||
|
func New(topN int) *Metrics {
|
||||||
|
byStatus := []string{"status_class", "action"}
|
||||||
|
byFile := []string{"file"}
|
||||||
|
|
||||||
|
m := &Metrics{
|
||||||
|
registry: prometheus.NewRegistry(),
|
||||||
|
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_requests_in_flight",
|
||||||
|
Help: "Requests under way.",
|
||||||
|
}),
|
||||||
|
requests: counterVec("smallwebwaf_requests_total",
|
||||||
|
"Requests, by the class of their status and their action.", byStatus),
|
||||||
|
requestBytes: counterVec("smallwebwaf_request_bytes_total",
|
||||||
|
"Request body bytes, by the class of the status and the action.",
|
||||||
|
byStatus),
|
||||||
|
responseBytes: counterVec("smallwebwaf_response_bytes_total",
|
||||||
|
"Response body bytes, by the class of the status and the action.",
|
||||||
|
byStatus),
|
||||||
|
requestDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
|
||||||
|
Name: "smallwebwaf_request_duration_seconds",
|
||||||
|
Help: "How long requests took, from their arrival to their end.",
|
||||||
|
}),
|
||||||
|
upstreamDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
|
||||||
|
Name: "smallwebwaf_upstream_duration_seconds",
|
||||||
|
Help: "How long requests passed to the app took, from then to their end.",
|
||||||
|
}),
|
||||||
|
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
|
||||||
|
"Requests that broke a rate limit, by its window.",
|
||||||
|
[]string{"window"}),
|
||||||
|
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
|
||||||
|
"Requests that passed a size or time limit, by its setting.",
|
||||||
|
[]string{"limit"}),
|
||||||
|
offences: counterVec("smallwebwaf_offences_total",
|
||||||
|
"Offences, by kind.", []string{"kind"}),
|
||||||
|
countries: newCountries(topN),
|
||||||
|
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_geojs_requests_total",
|
||||||
|
Help: "Requests to GeoJS.",
|
||||||
|
}),
|
||||||
|
GeoJSFailures: prometheus.NewCounter(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_geojs_failures_total",
|
||||||
|
Help: "Requests to GeoJS that failed.",
|
||||||
|
}),
|
||||||
|
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_geojs_unanswered_total",
|
||||||
|
Help: "Requests whose client counted as coming from an unknown " +
|
||||||
|
"country because GeoJS had not answered about it in time.",
|
||||||
|
}),
|
||||||
|
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
|
||||||
|
"Writes of each state file.", byFile),
|
||||||
|
stateFileWriteFailures: counterVec("smallwebwaf_state_file_write_failures_total",
|
||||||
|
"Writes of each state file that failed.", byFile),
|
||||||
|
stateFileLastWrite: gaugeVec("smallwebwaf_state_file_last_write_timestamp_seconds",
|
||||||
|
"When each state file was last written, in seconds since 1970.", byFile),
|
||||||
|
stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes",
|
||||||
|
"The size of each state file, as it was last written.", byFile),
|
||||||
|
stateFileEditsTakenIn: counterVec("smallwebwaf_state_file_edits_taken_in_total",
|
||||||
|
"Edits of each state file taken in while running.", byFile),
|
||||||
|
stateFileEditsSetAside: counterVec("smallwebwaf_state_file_edits_set_aside_total",
|
||||||
|
"Edits of each state file renamed to <name>.bad because they did not parse.",
|
||||||
|
byFile),
|
||||||
|
}
|
||||||
|
|
||||||
|
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
|
||||||
|
|
||||||
|
m.registry.MustRegister(
|
||||||
|
collectors.NewGoCollector(),
|
||||||
|
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
|
||||||
|
m.inFlight, m.requests, m.requestBytes, m.responseBytes,
|
||||||
|
m.requestDuration, m.upstreamDuration,
|
||||||
|
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
|
||||||
|
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
|
||||||
|
m.countries.refused,
|
||||||
|
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
|
||||||
|
m.stateFileWrites, m.stateFileWriteFailures,
|
||||||
|
m.stateFileLastWrite, m.stateFileSize,
|
||||||
|
m.stateFileEditsTakenIn, m.stateFileEditsSetAside,
|
||||||
|
)
|
||||||
|
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddBansAndClients adds the metrics read from the ledger and the table
|
||||||
|
// of clients as the metrics are asked for: the bans made since the start,
|
||||||
|
// by cause, the bans active and permanent at now, and the clients in the
|
||||||
|
// table.
|
||||||
|
func (m *Metrics) AddBansAndClients(
|
||||||
|
ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time,
|
||||||
|
) {
|
||||||
|
for _, cause := range []string{bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin} {
|
||||||
|
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_bans_made_total",
|
||||||
|
Help: "Bans made, by cause.",
|
||||||
|
ConstLabels: prometheus.Labels{"cause": cause},
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(ledger.Made(cause))
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
|
m.registry.MustRegister(
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_active_bans",
|
||||||
|
Help: "Bans active now, the permanent ones included.",
|
||||||
|
}, func() float64 {
|
||||||
|
active, _ := ledger.Count(now())
|
||||||
|
|
||||||
|
return float64(active)
|
||||||
|
}),
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_permanent_bans",
|
||||||
|
Help: "Permanent bans not lifted.",
|
||||||
|
}, func() float64 {
|
||||||
|
_, permanent := ledger.Count(now())
|
||||||
|
|
||||||
|
return float64(permanent)
|
||||||
|
}),
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_tracked_clients",
|
||||||
|
Help: "Clients in the table of clients.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(limiter.Len())
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddRules adds the metrics of the rule files: the requests that matched
|
||||||
|
// each rule, which RuleMatched counts, and the rules loaded from
|
||||||
|
// ruleFiles, read as the metrics are asked for. It is called once, before
|
||||||
|
// RuleMatched.
|
||||||
|
func (m *Metrics) AddRules(ruleFiles *rules.Files) {
|
||||||
|
m.ruleMatches = counterVec("smallwebwaf_rule_matches_total",
|
||||||
|
"Requests that matched a rule of the rule files, by its id and action.",
|
||||||
|
[]string{"rule_id", "action"})
|
||||||
|
|
||||||
|
m.registry.MustRegister(m.ruleMatches,
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_rules_loaded",
|
||||||
|
Help: "Rules loaded from the rule files.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(ruleFiles.Len())
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddRemoteLog adds the metrics of sending the log lines to
|
||||||
|
// SWWAF_LOG_REMOTE_URL, read from remote as the metrics are asked for: the
|
||||||
|
// lines sent, those dropped, and those waiting in the buffer.
|
||||||
|
func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
|
||||||
|
m.registry.MustRegister(
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_remote_log_lines_sent_total",
|
||||||
|
Help: "Log lines sent to SWWAF_LOG_REMOTE_URL.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(remote.Sent())
|
||||||
|
}),
|
||||||
|
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||||
|
Name: "smallwebwaf_remote_log_lines_dropped_total",
|
||||||
|
Help: "Log lines dropped: the oldest in a full buffer, and those " +
|
||||||
|
"whose sending failed.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(remote.Dropped())
|
||||||
|
}),
|
||||||
|
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||||
|
Name: "smallwebwaf_remote_log_buffer_depth",
|
||||||
|
Help: "Log lines in the buffer, waiting to be sent.",
|
||||||
|
}, func() float64 {
|
||||||
|
return float64(remote.Depth())
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServeHTTP answers with the metrics in the Prometheus text format.
|
||||||
|
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
|
m.handler.ServeHTTP(w, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RequestStarted counts a request as under way.
|
||||||
|
func (m *Metrics) RequestStarted() {
|
||||||
|
m.inFlight.Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RequestEnded counts a request that has ended, from its log line. limit
|
||||||
|
// is the setting whose size or time limit the request passed, "" if none.
|
||||||
|
// duration is how long the request took, and upstreamDuration how long it
|
||||||
|
// took from when it was passed to the app, zero if it was not.
|
||||||
|
func (m *Metrics) RequestEnded(
|
||||||
|
line *requestlog.Line, limit string, duration, upstreamDuration time.Duration,
|
||||||
|
) {
|
||||||
|
m.inFlight.Dec()
|
||||||
|
|
||||||
|
class := statusClass(line.Status)
|
||||||
|
m.requests.WithLabelValues(class, line.Action).Inc()
|
||||||
|
m.requestBytes.WithLabelValues(class, line.Action).Add(float64(line.RequestBytes))
|
||||||
|
m.responseBytes.WithLabelValues(class, line.Action).Add(float64(line.ResponseBytes))
|
||||||
|
m.requestDuration.Observe(duration.Seconds())
|
||||||
|
|
||||||
|
if upstreamDuration > 0 {
|
||||||
|
m.upstreamDuration.Observe(upstreamDuration.Seconds())
|
||||||
|
}
|
||||||
|
|
||||||
|
if line.LimitHit != "" {
|
||||||
|
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
if limit != "" {
|
||||||
|
m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
if line.Offence != "" {
|
||||||
|
m.offences.WithLabelValues(line.Offence).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
if line.Country != "" {
|
||||||
|
m.countries.add(line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RuleMatched counts a request that matched the rule id, whose action is
|
||||||
|
// action.
|
||||||
|
func (m *Metrics) RuleMatched(id, action string) {
|
||||||
|
m.ruleMatches.WithLabelValues(id, action).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// StateFileWritten counts a write of the state file name, of size bytes,
|
||||||
|
// that ended with err.
|
||||||
|
func (m *Metrics) StateFileWritten(name string, size int, err error) {
|
||||||
|
m.stateFileWrites.WithLabelValues(name).Inc()
|
||||||
|
|
||||||
|
// The series of failures is there from the first write, at zero until
|
||||||
|
// one fails.
|
||||||
|
failures := m.stateFileWriteFailures.WithLabelValues(name)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
failures.Inc()
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
m.stateFileLastWrite.WithLabelValues(name).SetToCurrentTime()
|
||||||
|
m.stateFileSize.WithLabelValues(name).Set(float64(size))
|
||||||
|
}
|
||||||
|
|
||||||
|
// StateFileEditTakenIn counts an admin's edit of the state file name
|
||||||
|
// taken in while smallwebwaf runs.
|
||||||
|
func (m *Metrics) StateFileEditTakenIn(name string) {
|
||||||
|
m.stateFileEditsTakenIn.WithLabelValues(name).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// StateFileEditSetAside counts an admin's edit of the state file name
|
||||||
|
// renamed to name.bad because it did not parse.
|
||||||
|
func (m *Metrics) StateFileEditSetAside(name string) {
|
||||||
|
m.stateFileEditsSetAside.WithLabelValues(name).Inc()
|
||||||
|
}
|
||||||
|
|
||||||
|
// statusClass returns the class of status, such as 2xx, or none when no
|
||||||
|
// status was sent.
|
||||||
|
func statusClass(status int) string {
|
||||||
|
if status == 0 {
|
||||||
|
return "none"
|
||||||
|
}
|
||||||
|
|
||||||
|
// A status's class is its hundreds: 404 is in 4xx.
|
||||||
|
const hundred = 100
|
||||||
|
|
||||||
|
return strconv.Itoa(status/hundred) + "xx"
|
||||||
|
}
|
||||||
|
|
||||||
|
// counterVec returns a counter named name, described by help, with a
|
||||||
|
// series for each set of values of labels.
|
||||||
|
func counterVec(name, help string, labels []string) *prometheus.CounterVec {
|
||||||
|
return prometheus.NewCounterVec(prometheus.CounterOpts{Name: name, Help: help},
|
||||||
|
labels)
|
||||||
|
}
|
||||||
|
|
||||||
|
// gaugeVec returns a gauge named name, described by help, with a series
|
||||||
|
// for each set of values of labels.
|
||||||
|
func gaugeVec(name, help string, labels []string) *prometheus.GaugeVec {
|
||||||
|
return prometheus.NewGaugeVec(prometheus.GaugeOpts{Name: name, Help: help}, labels)
|
||||||
|
}
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/subtle"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// answerAdmin answers a request for smallwebwaf itself, under
|
||||||
|
// /_smallwebwaf/, once it has passed the checks: GET MetricsPath with
|
||||||
|
// SWWAF_METRICS_TOKEN gets the metrics, and without it is refused with
|
||||||
|
// 401. Any other request gets 404, as the metrics do while
|
||||||
|
// SWWAF_METRICS_TOKEN is unset.
|
||||||
|
func (rq *request) answerAdmin() {
|
||||||
|
rq.line.Action = requestlog.ActionAdmin
|
||||||
|
rq.startClientResponseTimeout()
|
||||||
|
|
||||||
|
token := rq.h.config.MetricsToken
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case token == "" || rq.in.Method != http.MethodGet || rq.in.URL.Path != MetricsPath:
|
||||||
|
http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound)
|
||||||
|
case !hasToken(rq.in, token):
|
||||||
|
rq.out.Header().Set("WWW-Authenticate", "Bearer")
|
||||||
|
rq.answer(refusal{
|
||||||
|
status: http.StatusUnauthorized,
|
||||||
|
action: requestlog.ActionAdmin,
|
||||||
|
})
|
||||||
|
default:
|
||||||
|
rq.h.metrics.ServeHTTP(rq.out, rq.in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// hasToken reports whether r carries token, as Authorization: Bearer
|
||||||
|
// <token>.
|
||||||
|
func hasToken(r *http.Request, token string) bool {
|
||||||
|
scheme, sent, _ := strings.Cut(r.Header.Get("Authorization"), " ")
|
||||||
|
|
||||||
|
return strings.EqualFold(scheme, "Bearer") &&
|
||||||
|
subtle.ConstantTimeCompare([]byte(sent), []byte(token)) == 1
|
||||||
|
}
|
||||||
@@ -0,0 +1,125 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 covers the
|
||||||
|
// request at now, and notes for the log line when that ban ends.
|
||||||
|
func (rq *request) banned(now time.Time) bool {
|
||||||
|
check := rq.h.ledger.Check
|
||||||
|
if rq.h.config.Observe {
|
||||||
|
check = rq.h.ledger.Find // in observe mode the ban refuses nothing
|
||||||
|
}
|
||||||
|
|
||||||
|
ban, banned := check(rq.client, now)
|
||||||
|
if banned {
|
||||||
|
rq.line.BanExpires = banExpires(ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
return banned
|
||||||
|
}
|
||||||
|
|
||||||
|
// limitBroken counts the request for the rate limits at now, notes the
|
||||||
|
// client's counts for the log line, and reports whether the request takes
|
||||||
|
// the client over a limit. In enforce mode such a request bans the
|
||||||
|
// client's netblock, and sets the client's counters back to zero; in
|
||||||
|
// observe mode it does neither.
|
||||||
|
func (rq *request) limitBroken(now time.Time) bool {
|
||||||
|
group := clientGroup(rq.client)
|
||||||
|
|
||||||
|
counts, hit, over := rq.h.limiter.Count(group, now)
|
||||||
|
rq.line.Counts = counts
|
||||||
|
|
||||||
|
if !over {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
rq.line.LimitHit = hit.Window
|
||||||
|
rq.line.Offence = requestlog.OffenceLimit
|
||||||
|
|
||||||
|
if rq.h.config.Observe {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
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: rq.noted(now),
|
||||||
|
Requests: rq.netblockRequests(netblock),
|
||||||
|
})
|
||||||
|
rq.h.limiter.Reset(group)
|
||||||
|
rq.line.BanExpires = banExpires(ban)
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// banForAttack bans the client's netblock at now for a clear sign of
|
||||||
|
// attack, the match of rule, a ban rule.
|
||||||
|
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
|
||||||
|
netblock := rq.netblock()
|
||||||
|
ban := rq.h.ledger.BanForAttack(netblock, now, bans.Notes{
|
||||||
|
Country: rq.line.Country,
|
||||||
|
RuleID: rule.ID,
|
||||||
|
Target: rule.Target,
|
||||||
|
Request: rq.noted(now),
|
||||||
|
Requests: rq.netblockRequests(netblock),
|
||||||
|
})
|
||||||
|
rq.line.BanExpires = banExpires(ban)
|
||||||
|
}
|
||||||
|
|
||||||
|
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, as the
|
||||||
|
// notes of the ban it makes keep it.
|
||||||
|
func (rq *request) noted(now time.Time) bans.Request {
|
||||||
|
return 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(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// netblockRequests is how many requests netblock has sent since it was
|
||||||
|
// first seen, this one included: the histories count it only once it has
|
||||||
|
// ended.
|
||||||
|
func (rq *request) netblockRequests(netblock netip.Prefix) int64 {
|
||||||
|
return rq.h.limiter.Requests(netblock) + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,457 @@
|
|||||||
|
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),
|
||||||
|
Cause: bans.CauseLimit,
|
||||||
|
Reason: "requests per minute over the limit of 1",
|
||||||
|
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: bans.EarlierBans{},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
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 != (bans.EarlierBans{Limit: 1}) {
|
||||||
|
t.Errorf("bans %+v, want two, the second with one earlier ban for a limit", 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()
|
||||||
|
|
||||||
|
line, _ := s.requestWithHeader(from, path, "", status, action)
|
||||||
|
|
||||||
|
return line
|
||||||
|
}
|
||||||
|
|
||||||
|
// requestWithHeader is request with header, such as "Authorization:
|
||||||
|
// Bearer x", added to the request unless it is "". It returns the body of
|
||||||
|
// the answer too.
|
||||||
|
func (s *sender) requestWithHeader(
|
||||||
|
from, path, header string, status int, action string,
|
||||||
|
) (logLine, string) {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
if header != "" {
|
||||||
|
header += "\r\n"
|
||||||
|
}
|
||||||
|
|
||||||
|
conn := dial(s.t, s.addr)
|
||||||
|
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+
|
||||||
|
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+
|
||||||
|
header+"\r\n")
|
||||||
|
|
||||||
|
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
|
||||||
|
if err != nil {
|
||||||
|
s.t.Fatalf("set read deadline: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var got answer
|
||||||
|
|
||||||
|
res, err := http.ReadResponse(bufio.NewReader(conn), nil)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case err == nil:
|
||||||
|
got = readAnswer(res)
|
||||||
|
case !errors.Is(err, io.ErrUnexpectedEOF):
|
||||||
|
s.t.Fatalf("read response: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
if got.status != status {
|
||||||
|
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from,
|
||||||
|
got.status, status)
|
||||||
|
}
|
||||||
|
|
||||||
|
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
|
||||||
|
s.sent++
|
||||||
|
wantLine(s.t, line, status, action)
|
||||||
|
|
||||||
|
return line, string(got.body)
|
||||||
|
}
|
||||||
@@ -46,6 +46,7 @@ func (b *requestBody) Read(p []byte) (int, error) {
|
|||||||
b.rq.refuse(refusal{
|
b.rq.refuse(refusal{
|
||||||
status: http.StatusRequestEntityTooLarge,
|
status: http.StatusRequestEntityTooLarge,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
|
limit: "SWWAF_REQUEST_MAX_BYTES",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -81,6 +82,7 @@ func (b *responseBody) Read(p []byte) (int, error) {
|
|||||||
b.rq.refuse(refusal{
|
b.rq.refuse(refusal{
|
||||||
status: http.StatusBadGateway,
|
status: http.StatusBadGateway,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
|
limit: "SWWAF_RESPONSE_MAX_BYTES",
|
||||||
})
|
})
|
||||||
|
|
||||||
return n, errResponseTooLarge
|
return n, errResponseTooLarge
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package proxy
|
package proxy
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
@@ -48,6 +49,33 @@ func clientAddress(
|
|||||||
return client
|
return client
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// requestIDHeader carries the request's id, from traefik and to the app.
|
||||||
|
const requestIDHeader = "X-Request-ID"
|
||||||
|
|
||||||
|
// requestID is the request's id: the one a trusted proxy sent, or a new
|
||||||
|
// random one. A peer outside the trusted proxies did not come through
|
||||||
|
// traefik, so the id it sends is its own claim, and is replaced.
|
||||||
|
func requestID(r *http.Request, peerTrusted bool) string {
|
||||||
|
id := r.Header.Get(requestIDHeader)
|
||||||
|
if !peerTrusted || id == "" {
|
||||||
|
id = rand.Text()
|
||||||
|
}
|
||||||
|
|
||||||
|
return id
|
||||||
|
}
|
||||||
|
|
||||||
|
// scheme is how the client reached traefik, as a trusted proxy says in
|
||||||
|
// X-Forwarded-Proto, or otherwise http, the only scheme smallwebwaf
|
||||||
|
// serves.
|
||||||
|
func scheme(r *http.Request, peerTrusted bool) string {
|
||||||
|
proto := r.Header.Get("X-Forwarded-Proto")
|
||||||
|
if !peerTrusted || proto == "" {
|
||||||
|
return "http"
|
||||||
|
}
|
||||||
|
|
||||||
|
return proto
|
||||||
|
}
|
||||||
|
|
||||||
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
|
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
|
||||||
const ipv6GroupPrefix = 64
|
const ipv6GroupPrefix = 64
|
||||||
|
|
||||||
|
|||||||
@@ -14,10 +14,14 @@ const (
|
|||||||
appHost = "app.example"
|
appHost = "app.example"
|
||||||
// client is the client's address, as a proxy names it.
|
// client is the client's address, as a proxy names it.
|
||||||
client = "203.0.113.9"
|
client = "203.0.113.9"
|
||||||
// forwardedFor is the header that lists the client and its proxies.
|
// forwardedFor is the header that lists the client and its proxies,
|
||||||
forwardedFor = "X-Forwarded-For"
|
// and forwardedProto the one that gives the scheme the client used.
|
||||||
// secure is the scheme a client reached traefik with.
|
forwardedFor = "X-Forwarded-For"
|
||||||
|
forwardedProto = "X-Forwarded-Proto"
|
||||||
|
// secure is the scheme a client reached traefik with, and plain the
|
||||||
|
// one smallwebwaf serves.
|
||||||
secure = "https"
|
secure = "https"
|
||||||
|
plain = "http"
|
||||||
)
|
)
|
||||||
|
|
||||||
// appHeaders is what the app tells about the headers it received.
|
// appHeaders is what the app tells about the headers it received.
|
||||||
@@ -65,13 +69,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
|
|||||||
func clientAddressCases() []clientAddressCase {
|
func clientAddressCases() []clientAddressCase {
|
||||||
trusted := map[string]string{trustedProxies: trustLocalhost}
|
trusted := map[string]string{trustedProxies: trustLocalhost}
|
||||||
forged := http.Header{
|
forged := http.Header{
|
||||||
forwardedFor: {client},
|
forwardedFor: {client},
|
||||||
"X-Forwarded-Host": {"forged.example"},
|
"X-Forwarded-Host": {"forged.example"},
|
||||||
"X-Forwarded-Proto": {secure},
|
forwardedProto: {secure},
|
||||||
"X-Real-Ip": {client},
|
"X-Real-Ip": {client},
|
||||||
}
|
}
|
||||||
replaced := appHeaders{
|
replaced := appHeaders{
|
||||||
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http",
|
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
|
||||||
}
|
}
|
||||||
|
|
||||||
return []clientAddressCase{{
|
return []clientAddressCase{{
|
||||||
@@ -87,10 +91,10 @@ func clientAddressCases() []clientAddressCase {
|
|||||||
"outside the trusted proxies from the right",
|
"outside the trusted proxies from the right",
|
||||||
env: trusted,
|
env: trusted,
|
||||||
header: http.Header{
|
header: http.Header{
|
||||||
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
|
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
|
||||||
"X-Forwarded-Host": {appHost},
|
"X-Forwarded-Host": {appHost},
|
||||||
"X-Forwarded-Proto": {secure},
|
forwardedProto: {secure},
|
||||||
"X-Real-Ip": {client},
|
"X-Real-Ip": {client},
|
||||||
},
|
},
|
||||||
wantClient: client,
|
wantClient: client,
|
||||||
wantApp: appHeaders{
|
wantApp: appHeaders{
|
||||||
@@ -138,7 +142,7 @@ func requestWithHeaders(
|
|||||||
Host: r.Host,
|
Host: r.Host,
|
||||||
ForwardedFor: r.Header.Get(forwardedFor),
|
ForwardedFor: r.Header.Get(forwardedFor),
|
||||||
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
|
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
|
||||||
ForwardedProto: r.Header.Get("X-Forwarded-Proto"),
|
ForwardedProto: r.Header.Get(forwardedProto),
|
||||||
RealIP: r.Header.Get("X-Real-IP"),
|
RealIP: r.Header.Get("X-Real-IP"),
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -21,14 +21,18 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
|
|||||||
// the last one would have it refused.
|
// the last one would have it refused.
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"})
|
addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"})
|
||||||
|
|
||||||
const healthChecks = 3
|
const (
|
||||||
|
healthChecks = 3
|
||||||
|
contentType = "text/plain; charset=utf-8"
|
||||||
|
)
|
||||||
|
|
||||||
for range healthChecks {
|
for range healthChecks {
|
||||||
got := get(t, addr, proxy.HealthPath)
|
got := get(t, addr, proxy.HealthPath)
|
||||||
wantStatus(t, got, http.StatusOK)
|
wantStatus(t, got, http.StatusOK)
|
||||||
|
|
||||||
if string(got.body) != "ok\n" {
|
if string(got.body) != "ok\n" || got.header.Get("Content-Type") != contentType {
|
||||||
t.Errorf("health endpoint answered %q, want ok", got.body)
|
t.Errorf("health endpoint answered %q with Content-Type %q, want ok "+
|
||||||
|
"with %q", got.body, got.header.Get("Content-Type"), contentType)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -37,6 +41,11 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
|
|||||||
lines := out.requestLines(t, healthChecks+1)
|
lines := out.requestLines(t, healthChecks+1)
|
||||||
for _, line := range lines[:healthChecks] {
|
for _, line := range lines[:healthChecks] {
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionAdmin)
|
wantLine(t, line, http.StatusOK, requestlog.ActionAdmin)
|
||||||
|
|
||||||
|
if line.ResponseContentType != contentType {
|
||||||
|
t.Errorf("health check's log line has response_content_type %q, "+
|
||||||
|
"want %q", line.ResponseContentType, contentType)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward)
|
wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward)
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestForSmallwebwafIsRefusedOnlyWithoutTheToken(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now,
|
||||||
|
map[string]string{metricsToken: token})
|
||||||
|
|
||||||
|
// The metrics and the 404 are neither forwarded nor refused; the 401
|
||||||
|
// is refused.
|
||||||
|
scrape(t, addr)
|
||||||
|
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
|
||||||
|
wantStatus(t, get(t, addr, proxy.MetricsPath), http.StatusUnauthorized)
|
||||||
|
out.requestLines(t, 3)
|
||||||
|
|
||||||
|
history := historyOf(t, server, localhost)
|
||||||
|
if history.Requests != 3 || history.Forwarded != 0 || history.Refused != 1 {
|
||||||
|
t.Errorf("history counts %d requests, %d forwarded and %d refused, "+
|
||||||
|
"want 3, 0 and 1", history.Requests, history.Forwarded, history.Refused)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// historyOf returns the history of the client at addr.
|
||||||
|
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
client := netip.MustParsePrefix(addr + "/32")
|
||||||
|
for _, c := range server.Limiter.Snapshot() {
|
||||||
|
if c.Client == client {
|
||||||
|
return c.History
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Fatalf("%s is not in the table", client)
|
||||||
|
|
||||||
|
return ratelimit.History{}
|
||||||
|
}
|
||||||
@@ -55,6 +55,7 @@ func TestRequestBodyLimit(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
requestMaxBytes: sizeLimitSetting,
|
requestMaxBytes: sizeLimitSetting,
|
||||||
|
metricsToken: token,
|
||||||
})
|
})
|
||||||
|
|
||||||
var body io.Reader = bytes.NewReader(make([]byte, tc.size))
|
var body io.Reader = bytes.NewReader(make([]byte, tc.size))
|
||||||
@@ -66,6 +67,13 @@ func TestRequestBodyLimit(t *testing.T) {
|
|||||||
tc.want)
|
tc.want)
|
||||||
wantLine(t, out.requestLine(t), tc.want, tc.action)
|
wantLine(t, out.requestLine(t), tc.want, tc.action)
|
||||||
|
|
||||||
|
hits := 0
|
||||||
|
if tc.action == requestlog.ActionTooLarge {
|
||||||
|
hits = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
wantLimitHits(t, addr, requestMaxBytes, hits)
|
||||||
|
|
||||||
if tc.refusedBeforeApp && calls.Load() != 0 {
|
if tc.refusedBeforeApp && calls.Load() != 0 {
|
||||||
t.Errorf("the app was called %d times, want never", calls.Load())
|
t.Errorf("the app was called %d times, want never", calls.Load())
|
||||||
}
|
}
|
||||||
@@ -106,6 +114,7 @@ func TestResponseBodyLimit(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
responseMaxBytes: sizeLimitSetting,
|
responseMaxBytes: sizeLimitSetting,
|
||||||
|
metricsToken: token,
|
||||||
})
|
})
|
||||||
|
|
||||||
got := get(t, addr, "/download")
|
got := get(t, addr, "/download")
|
||||||
@@ -123,6 +132,13 @@ func TestResponseBodyLimit(t *testing.T) {
|
|||||||
if line.UpstreamStatus != http.StatusOK {
|
if line.UpstreamStatus != http.StatusOK {
|
||||||
t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
|
t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
hits := 0
|
||||||
|
if tc.action == requestlog.ActionTooLarge {
|
||||||
|
hits = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
wantLimitHits(t, addr, responseMaxBytes, hits)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,479 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/netip"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
|
||||||
|
metricsTopN = "SWWAF_METRICS_TOP_N"
|
||||||
|
// token is the SWWAF_METRICS_TOKEN the tests set, and bearer how a
|
||||||
|
// request carries it.
|
||||||
|
token = "0123456789abcdef0123456789abcdef"
|
||||||
|
bearer = "Bearer " + token
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMetricsAreOffWhileTheTokenIsUnset(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var calls atomic.Int32
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {
|
||||||
|
calls.Add(1)
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, nil)
|
||||||
|
|
||||||
|
// An empty token does not match the unset one either.
|
||||||
|
for i, authorization := range []string{bearer, "Bearer ", ""} {
|
||||||
|
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
|
||||||
|
if authorization != "" {
|
||||||
|
req.Header.Set("Authorization", authorization)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantStatus(t, do(t, req), http.StatusNotFound)
|
||||||
|
wantLine(t, out.requestLines(t, i+1)[i], http.StatusNotFound,
|
||||||
|
requestlog.ActionAdmin)
|
||||||
|
}
|
||||||
|
|
||||||
|
if calls.Load() != 0 {
|
||||||
|
t.Errorf("the app was called %d times, want never", calls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsNeedTheTokenAndOtherPathsAreNotFound(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
var calls atomic.Int32
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {
|
||||||
|
calls.Add(1)
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
|
||||||
|
|
||||||
|
for i, tc := range []struct {
|
||||||
|
method, path, authorization string
|
||||||
|
status int
|
||||||
|
}{
|
||||||
|
{http.MethodGet, proxy.MetricsPath, "", http.StatusUnauthorized},
|
||||||
|
{
|
||||||
|
http.MethodGet, proxy.MetricsPath, "Bearer " + strings.ToUpper(token),
|
||||||
|
http.StatusUnauthorized,
|
||||||
|
},
|
||||||
|
{http.MethodGet, proxy.MetricsPath, "Basic " + token, http.StatusUnauthorized},
|
||||||
|
{http.MethodGet, proxy.MetricsPath, bearer, http.StatusOK},
|
||||||
|
{http.MethodGet, proxy.MetricsPath, "bearer " + token, http.StatusOK},
|
||||||
|
{http.MethodPost, proxy.MetricsPath, bearer, http.StatusNotFound},
|
||||||
|
{http.MethodGet, proxy.MetricsPath + "/", bearer, http.StatusNotFound},
|
||||||
|
{http.MethodGet, "/_smallwebwaf/bans", bearer, http.StatusNotFound},
|
||||||
|
{http.MethodPost, proxy.HealthPath, "", http.StatusNotFound},
|
||||||
|
} {
|
||||||
|
req := newRequest(t, tc.method, addr, tc.path, http.NoBody)
|
||||||
|
if tc.authorization != "" {
|
||||||
|
req.Header.Set("Authorization", tc.authorization)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := do(t, req)
|
||||||
|
wantStatus(t, got, tc.status)
|
||||||
|
wantLine(t, out.requestLines(t, i+1)[i], tc.status, requestlog.ActionAdmin)
|
||||||
|
|
||||||
|
if tc.status == http.StatusUnauthorized &&
|
||||||
|
got.header.Get("WWW-Authenticate") != "Bearer" {
|
||||||
|
t.Errorf("%q was answered without WWW-Authenticate: Bearer",
|
||||||
|
tc.authorization)
|
||||||
|
}
|
||||||
|
|
||||||
|
if tc.status == http.StatusOK &&
|
||||||
|
!strings.Contains(string(got.body), "# TYPE smallwebwaf_requests_total counter") {
|
||||||
|
t.Errorf("the metrics are\n%s", got.body)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if calls.Load() != 0 {
|
||||||
|
t.Errorf("the app was called %d times, want never", calls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsAreAskedForThroughTheChecks(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithClock(t, "", map[string]string{
|
||||||
|
metricsToken: token,
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Asking for the metrics counts toward the client's limit of one
|
||||||
|
// request a minute, so its next request breaks it, and bans it. A
|
||||||
|
// banned client is refused the metrics too.
|
||||||
|
s.scrape(client)
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
s.requestWithHeader(client, proxy.MetricsPath, "Authorization: "+bearer,
|
||||||
|
http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsCountTheTraffic(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
arrived, release := make(chan struct{}), make(chan struct{})
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
|
||||||
|
if r.URL.Path == "/held" {
|
||||||
|
close(arrived)
|
||||||
|
<-release
|
||||||
|
}
|
||||||
|
|
||||||
|
_, _ = io.WriteString(w, "hello")
|
||||||
|
})
|
||||||
|
releaseApp := sync.OnceFunc(func() { close(release) })
|
||||||
|
t.Cleanup(releaseApp)
|
||||||
|
|
||||||
|
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
|
||||||
|
|
||||||
|
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
|
||||||
|
wantStatus(t, got, http.StatusOK)
|
||||||
|
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
|
||||||
|
out.requestLines(t, 2)
|
||||||
|
|
||||||
|
forward := `{action="forward",status_class="2xx"}`
|
||||||
|
notFound := `{action="admin",status_class="4xx"}`
|
||||||
|
|
||||||
|
// The request for the metrics is itself under way.
|
||||||
|
metrics := scrape(t, addr)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_requests_total"+forward, 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_requests_total"+notFound, 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_request_bytes_total"+forward, 3)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
|
||||||
|
float64(len("Not Found\n")))
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
|
||||||
|
metric(t, metrics, "go_goroutines")
|
||||||
|
metric(t, metrics, "process_start_time_seconds")
|
||||||
|
|
||||||
|
// A request the app holds is under way until it ends.
|
||||||
|
httpClient := newClient(t)
|
||||||
|
held := newRequest(t, http.MethodGet, addr, "/held", http.NoBody)
|
||||||
|
ended := make(chan error, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
res, err := httpClient.Do(held)
|
||||||
|
if err == nil {
|
||||||
|
err = readAnswer(res).err
|
||||||
|
}
|
||||||
|
|
||||||
|
ended <- err
|
||||||
|
}()
|
||||||
|
|
||||||
|
<-arrived
|
||||||
|
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
|
||||||
|
releaseApp()
|
||||||
|
|
||||||
|
err := <-ended
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("held request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out.requestLines(t, 5)
|
||||||
|
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsCountLimitsAndBans(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||||
|
denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
)
|
||||||
|
|
||||||
|
s, clk, _ := startWithClock(t, "", map[string]string{
|
||||||
|
metricsToken: token,
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
rateLimitExemptNets: scraper,
|
||||||
|
denyNets: denied,
|
||||||
|
banResponse: "close",
|
||||||
|
limitBanDuration: "1h",
|
||||||
|
maxBanDuration: "2h",
|
||||||
|
})
|
||||||
|
|
||||||
|
// SWWAF_BAN_RESPONSE=close sends no status at all.
|
||||||
|
s.get(denied, 0, requestlog.ActionDenied)
|
||||||
|
|
||||||
|
// A first broken limit bans for an hour.
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(client, 0, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
metrics := s.scrape(scraper)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
|
||||||
|
|
||||||
|
clk.advance(time.Hour)
|
||||||
|
wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
|
||||||
|
|
||||||
|
// A limit broken again right after would ban for three hours, longer
|
||||||
|
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(client, 0, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
metrics = s.scrape(scraper)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
|
||||||
|
// denied, client, and the scraper as of its earlier requests.
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||||
|
|
||||||
|
s, clk, server := startWithClock(t, "", map[string]string{
|
||||||
|
metricsToken: token,
|
||||||
|
rateLimitExemptNets: scraper,
|
||||||
|
})
|
||||||
|
|
||||||
|
const admins = `smallwebwaf_bans_made_total{cause="admin"}`
|
||||||
|
|
||||||
|
wantMetric(t, s.scrape(scraper), admins, 0)
|
||||||
|
|
||||||
|
// As an admin's edit of bans.json that adds a ban is taken in.
|
||||||
|
server.Ledger.LoadEdit([]bans.Ban{{
|
||||||
|
Netblock: netip.MustParsePrefix(client + "/32"),
|
||||||
|
Start: clk.Now(),
|
||||||
|
}})
|
||||||
|
|
||||||
|
wantMetric(t, s.scrape(scraper), admins, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const fromFR = "198.51.100.20"
|
||||||
|
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
_, _ = io.WriteString(w, "hello")
|
||||||
|
})
|
||||||
|
env := map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
metricsToken: token,
|
||||||
|
metricsTopN: "2",
|
||||||
|
deniedCountries: "kp",
|
||||||
|
}
|
||||||
|
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
|
||||||
|
|
||||||
|
// The answers are kept before the requests, so that none waits for
|
||||||
|
// GeoJS.
|
||||||
|
server.GeoJS.Load([]lookup.Answer{
|
||||||
|
keptAnswer(fromKP, "KP"), keptAnswer(fromDE, "DE"), keptAnswer(fromFR, "FR"),
|
||||||
|
})
|
||||||
|
|
||||||
|
lines := 0
|
||||||
|
send := func(from string, times, status int) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for range times {
|
||||||
|
req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc"))
|
||||||
|
req.Header.Set(forwardedFor, from)
|
||||||
|
wantStatus(t, do(t, req), status)
|
||||||
|
|
||||||
|
// Each is counted before the next is sent, so that the
|
||||||
|
// countries are ranked in the order sent.
|
||||||
|
lines++
|
||||||
|
out.requestLines(t, lines)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// With two countries of their own, the third is counted as other.
|
||||||
|
send(fromKP, 3, http.StatusForbidden)
|
||||||
|
send(fromDE, 2, http.StatusOK)
|
||||||
|
send(fromFR, 1, http.StatusOK)
|
||||||
|
|
||||||
|
metrics := scrape(t, addr)
|
||||||
|
lines++
|
||||||
|
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
|
||||||
|
float64(3*len("Forbidden\n")))
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
|
||||||
|
float64(len("hello")))
|
||||||
|
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
|
||||||
|
|
||||||
|
// Once FR is busier than DE, it takes DE's place: its series counts
|
||||||
|
// from then on, and DE's is gone.
|
||||||
|
send(fromFR, 3, http.StatusOK)
|
||||||
|
|
||||||
|
metrics = scrape(t, addr)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
|
||||||
|
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
|
||||||
|
wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
geojs := httptest.NewServer(http.HandlerFunc(
|
||||||
|
func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
}))
|
||||||
|
t.Cleanup(geojs.Close)
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
metricsToken: token,
|
||||||
|
deniedCountries: "kp",
|
||||||
|
})
|
||||||
|
|
||||||
|
// GeoJS fails, so the client counts as coming from an unknown country,
|
||||||
|
// which SWWAF_DENIED_COUNTRIES does not refuse.
|
||||||
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
|
req.Header.Set(forwardedFor, fromDE)
|
||||||
|
wantStatus(t, do(t, req), http.StatusOK)
|
||||||
|
|
||||||
|
// The client stops waiting for GeoJS after a second, so GeoJS's
|
||||||
|
// failure can come after its request has ended.
|
||||||
|
deadline := time.Now().Add(waitLimit)
|
||||||
|
metrics := scrape(t, addr)
|
||||||
|
|
||||||
|
for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
|
||||||
|
time.Now().Before(deadline) {
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
|
||||||
|
metrics = scrape(t, addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// keptAnswer returns GeoJS's answer that the client at addr is in
|
||||||
|
// country, given now.
|
||||||
|
func keptAnswer(addr, country string) lookup.Answer {
|
||||||
|
now := time.Now()
|
||||||
|
|
||||||
|
return lookup.Answer{
|
||||||
|
Client: netip.MustParsePrefix(addr + "/32"), Country: country,
|
||||||
|
Answered: now, Used: now,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// scrape asks smallwebwaf at addr for the metrics, with the token, and
|
||||||
|
// returns them.
|
||||||
|
func scrape(t *testing.T, addr string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
|
||||||
|
req.Header.Set("Authorization", bearer)
|
||||||
|
|
||||||
|
got := do(t, req)
|
||||||
|
if got.status != http.StatusOK {
|
||||||
|
t.Fatalf("the metrics were answered %d", got.status)
|
||||||
|
}
|
||||||
|
|
||||||
|
return string(got.body)
|
||||||
|
}
|
||||||
|
|
||||||
|
// scrape asks for the metrics, with the token, from the client at from,
|
||||||
|
// and returns them.
|
||||||
|
func (s *sender) scrape(from string) string {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
|
_, metrics := s.requestWithHeader(from, proxy.MetricsPath, "Authorization: "+bearer,
|
||||||
|
http.StatusOK, requestlog.ActionAdmin)
|
||||||
|
|
||||||
|
return metrics
|
||||||
|
}
|
||||||
|
|
||||||
|
// metric returns the value of series in metrics, which are in the
|
||||||
|
// Prometheus text format. series is a name and its labels in the order of
|
||||||
|
// their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
|
||||||
|
// the test if there is no such series.
|
||||||
|
func metric(t *testing.T, metrics, series string) float64 {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for line := range strings.Lines(metrics) {
|
||||||
|
value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ")
|
||||||
|
if !found {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
number, err := strconv.ParseFloat(value, 64)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("%s has the value %q", series, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
return number
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Fatalf("no series %s in the metrics:\n%s", series, metrics)
|
||||||
|
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantMetric checks the value of series in metrics, as metric reads it.
|
||||||
|
func wantMetric(t *testing.T, metrics, series string, want float64) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := metric(t, metrics, series)
|
||||||
|
if got != want {
|
||||||
|
t.Errorf("%s is %v, want %v", series, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantNoSeries checks that metrics have no series series.
|
||||||
|
func wantNoSeries(t *testing.T, metrics, series string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if strings.Contains(metrics, "\n"+series+" ") {
|
||||||
|
t.Errorf("there is a series %s", series)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantLimitHits checks that the metrics of smallwebwaf at addr count hits
|
||||||
|
// requests that passed the size or time limit of the setting limit, with
|
||||||
|
// no series for it when hits is 0.
|
||||||
|
func wantLimitHits(t *testing.T, addr, limit string, hits int) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}`
|
||||||
|
metrics := scrape(t, addr)
|
||||||
|
|
||||||
|
if hits == 0 {
|
||||||
|
wantNoSeries(t, metrics, series)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
wantMetric(t, metrics, series, float64(hits))
|
||||||
|
}
|
||||||
@@ -0,0 +1,203 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// observe is the value of SWWAF_MODE for observe mode.
|
||||||
|
const observe = "observe"
|
||||||
|
|
||||||
|
func TestObserveModeForwardsWhatEnforceModeRefuses(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
banned = otherClient // under a ban read from bans.json
|
||||||
|
)
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
setting string // "" leaves SWWAF_MODE at its default
|
||||||
|
observe bool
|
||||||
|
}{
|
||||||
|
{"", false},
|
||||||
|
{"enforce", false},
|
||||||
|
{observe, true},
|
||||||
|
} {
|
||||||
|
t.Run(mode+"="+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[mode] = tc.setting
|
||||||
|
}
|
||||||
|
|
||||||
|
s, clk, server := startWithClock(t, geojsURL, env)
|
||||||
|
server.Ledger.Load([]bans.Ban{{
|
||||||
|
Netblock: netip.MustParsePrefix(banned + "/32"),
|
||||||
|
Start: clk.Now(),
|
||||||
|
Expires: clk.Now().Add(time.Hour),
|
||||||
|
}})
|
||||||
|
|
||||||
|
// fromDE's first request is within the limit of one a minute,
|
||||||
|
// and its second breaks it.
|
||||||
|
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
for _, sent := range []struct{ from, refusal string }{
|
||||||
|
{denied, requestlog.ActionDenied},
|
||||||
|
{banned, requestlog.ActionBanned},
|
||||||
|
{fromKP, requestlog.ActionCountryDenied},
|
||||||
|
{fromDE, requestlog.ActionRateLimited},
|
||||||
|
} {
|
||||||
|
if !tc.observe {
|
||||||
|
line := s.get(sent.from, http.StatusForbidden, sent.refusal)
|
||||||
|
wantWouldAction(t, line, "")
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Passed to the app, which answered it.
|
||||||
|
line := s.get(sent.from, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, sent.refusal)
|
||||||
|
|
||||||
|
if line.UpstreamStatus != http.StatusOK {
|
||||||
|
t.Errorf("log line has upstream_status %d, want 200",
|
||||||
|
line.UpstreamStatus)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server := startWithClock(t, "", map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
kept := bans.Ban{
|
||||||
|
Netblock: netip.MustParsePrefix(otherClient + "/32"),
|
||||||
|
Start: clk.Now(),
|
||||||
|
Expires: clk.Now().Add(time.Hour),
|
||||||
|
Cause: bans.CauseAdmin,
|
||||||
|
}
|
||||||
|
server.Ledger.Load([]bans.Ban{kept})
|
||||||
|
|
||||||
|
// No ban sets client's counters back to zero, so each request after
|
||||||
|
// the first breaks the limit of one a minute.
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
for range 2 {
|
||||||
|
line := s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit ||
|
||||||
|
line.BanExpires != "" {
|
||||||
|
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
|
||||||
|
"want minute, limit and none", line.LimitHit, line.Offence,
|
||||||
|
line.BanExpires)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The ban read from bans.json refuses nothing, and so counts no
|
||||||
|
// refusal in its notes, but is kept.
|
||||||
|
line := s.get(otherClient, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
if line.BanExpires != requestlog.FormatTime(kept.Expires) {
|
||||||
|
t.Errorf("log line has ban_expires %q, want %s", line.BanExpires,
|
||||||
|
requestlog.FormatTime(kept.Expires))
|
||||||
|
}
|
||||||
|
|
||||||
|
got := server.Ledger.Snapshot()
|
||||||
|
if len(got) != 1 || got[0] != kept {
|
||||||
|
t.Errorf("bans\n%+v\nwant only\n%+v", got, kept)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeKeepsTheSizeLimitsAndTheToken(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
|
||||||
|
var calls atomic.Int32
|
||||||
|
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
calls.Add(1)
|
||||||
|
answerWithSize(w, 2*sizeLimit, true)
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
|
mode: observe,
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
denyNets: denied,
|
||||||
|
requestMaxBytes: sizeLimitSetting,
|
||||||
|
responseMaxBytes: sizeLimitSetting,
|
||||||
|
metricsToken: token,
|
||||||
|
})
|
||||||
|
|
||||||
|
// SWWAF_DENY_NETS would refuse each request; instead a size limit or
|
||||||
|
// the missing token does.
|
||||||
|
for i, tc := range []struct {
|
||||||
|
method, path string
|
||||||
|
body io.Reader
|
||||||
|
status int
|
||||||
|
action string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
http.MethodPost, "/upload", bytes.NewReader(make([]byte, 2*sizeLimit)),
|
||||||
|
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
http.MethodGet, "/download", http.NoBody,
|
||||||
|
http.StatusBadGateway, requestlog.ActionTooLarge,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
http.MethodGet, proxy.MetricsPath, http.NoBody,
|
||||||
|
http.StatusUnauthorized, requestlog.ActionAdmin,
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
req := newRequest(t, tc.method, addr, tc.path, tc.body)
|
||||||
|
req.Header.Set(forwardedFor, denied)
|
||||||
|
wantStatus(t, do(t, req), tc.status)
|
||||||
|
|
||||||
|
line := out.requestLines(t, i+1)[i]
|
||||||
|
wantLine(t, line, tc.status, tc.action)
|
||||||
|
wantWouldAction(t, line, requestlog.ActionDenied)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The upload was refused before it reached the app.
|
||||||
|
if calls.Load() != 1 {
|
||||||
|
t.Errorf("the app was called %d times, want once", calls.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantWouldAction checks the request log line's would_action, and that a
|
||||||
|
// line that should have none has no such field.
|
||||||
|
func wantWouldAction(t *testing.T, line logLine, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got, present := line.fields["would_action"]
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case want == "" && present:
|
||||||
|
t.Errorf("log line has would_action %v, want none", got)
|
||||||
|
case want != "" && got != want:
|
||||||
|
t.Errorf("log line has would_action %v, want %s", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -14,6 +16,7 @@ import (
|
|||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -115,27 +118,34 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantRequestFields checks the log line's fields about the request.
|
// wantRequestFields checks the log line's fields about the request. Its
|
||||||
|
// time, its id and its timings are checked only for being there.
|
||||||
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
want := requestlog.Line{
|
hostname, _ := os.Hostname()
|
||||||
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
|
|
||||||
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery,
|
want := withTimings(line, requestlog.Line{
|
||||||
Protocol: "HTTP/1.1", Status: http.StatusTeapot,
|
Type: requestType, Time: line.Time, Instance: hostname,
|
||||||
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent),
|
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
|
||||||
|
Path: rawPath, Query: rawQuery, Protocol: protocol,
|
||||||
|
Status: http.StatusTeapot, RequestBytes: int64(sent),
|
||||||
ResponseBytes: int64(received), UserAgent: "test-agent",
|
ResponseBytes: int64(received), UserAgent: "test-agent",
|
||||||
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal,
|
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
|
||||||
DurationUpstreamTotal: line.DurationUpstreamTotal,
|
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
|
||||||
}
|
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
|
||||||
if line.Line != want {
|
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
||||||
|
})
|
||||||
|
if !reflect.DeepEqual(line.Line, want) {
|
||||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := time.Parse(time.RFC3339, line.Time)
|
_, err := time.Parse(time.RFC3339, line.Time)
|
||||||
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 {
|
if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
|
||||||
t.Errorf("log line has time %q and durations %v and %v",
|
line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
|
||||||
line.Time, line.DurationTotal, line.DurationUpstreamTotal)
|
t.Errorf("log line has time %q, request_id %q and durations %v and %v",
|
||||||
|
line.Time, line.RequestID, line.DurationTotal,
|
||||||
|
line.fields["duration_upstream_total"])
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -297,7 +307,7 @@ func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestServerHasTheFixedLimits(t *testing.T) {
|
func TestServerHasTheDefaultLimits(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
|
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
|
||||||
@@ -319,37 +329,48 @@ func TestServerHasTheFixedLimits(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRefusesHeadersOver32KiB(t *testing.T) {
|
func TestRefusesHeadersOverTheLimit(t *testing.T) {
|
||||||
t.Parallel()
|
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 {
|
for _, tc := range []struct {
|
||||||
size int
|
name string
|
||||||
want int
|
env map[string]string
|
||||||
|
limit int
|
||||||
}{
|
}{
|
||||||
{size: 32 << 10, want: http.StatusOK},
|
{"by default", nil, 32 << 10},
|
||||||
{size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge},
|
{"as set", map[string]string{clientHeaderMaxBytes: "8K"}, 8 << 10},
|
||||||
} {
|
} {
|
||||||
conn := dial(t, addr)
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end)
|
t.Parallel()
|
||||||
wantStatus(t, readResponse(t, conn), tc.want)
|
|
||||||
}
|
|
||||||
|
|
||||||
if calls.Load() != 1 {
|
var calls atomic.Int32
|
||||||
t.Errorf("the app was called %d times, want once", calls.Load())
|
|
||||||
|
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())
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -360,8 +381,13 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
|
|||||||
addr, out := startProxy(t, "http://"+localhost+":1", nil)
|
addr, out := startProxy(t, "http://"+localhost+":1", nil)
|
||||||
|
|
||||||
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
|
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
|
||||||
wantLine(t, out.requestLine(t), http.StatusBadGateway,
|
|
||||||
requestlog.ActionUpstreamError)
|
line := out.requestLine(t)
|
||||||
|
wantLine(t, line, http.StatusBadGateway, requestlog.ActionUpstreamError)
|
||||||
|
|
||||||
|
// There never was a connection to the app, nor an answer from it.
|
||||||
|
wantTimings(t, line, "duration_total", "duration_checks",
|
||||||
|
"duration_upstream_total")
|
||||||
|
|
||||||
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
|
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
|
||||||
return line["type"] == "process" && line["msg"] == "request to the app failed"
|
return line["type"] == "process" && line["msg"] == "request to the app failed"
|
||||||
|
|||||||
+103
-39
@@ -8,25 +8,16 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
|
||||||
// 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.
|
// How smallwebwaf keeps connections to the app open between requests.
|
||||||
@@ -35,10 +26,18 @@ const (
|
|||||||
appIdleConnTimeout = 90 * time.Second
|
appIdleConnTimeout = 90 * time.Second
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// adminPrefix starts the path of every request for smallwebwaf itself,
|
||||||
|
// which never reaches the app.
|
||||||
|
const adminPrefix = "/_smallwebwaf/"
|
||||||
|
|
||||||
// HealthPath is smallwebwaf's health endpoint, which the container's
|
// HealthPath is smallwebwaf's health endpoint, which the container's
|
||||||
// health check asks.
|
// health check asks.
|
||||||
const HealthPath = "/_smallwebwaf/healthz"
|
const HealthPath = "/_smallwebwaf/healthz"
|
||||||
|
|
||||||
|
// MetricsPath is where the metrics are, for a request that carries
|
||||||
|
// SWWAF_METRICS_TOKEN.
|
||||||
|
const MetricsPath = "/_smallwebwaf/metrics"
|
||||||
|
|
||||||
// Params are what New needs.
|
// Params are what New needs.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
@@ -49,39 +48,83 @@ type Params struct {
|
|||||||
// GeoJSURL is where clients' countries are looked up, normally
|
// GeoJSURL is where clients' countries are looked up, normally
|
||||||
// lookup.URL. GeoJS is asked only while a country list is set.
|
// lookup.URL. GeoJS is asked only while a country list is set.
|
||||||
GeoJSURL string
|
GeoJSURL string
|
||||||
|
// Now tells the time by which requests are counted for the rate
|
||||||
|
// 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
|
||||||
|
// Rules are the rule files' rules, which each request is checked
|
||||||
|
// against.
|
||||||
|
Rules *rules.Files
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server is the server smallwebwaf runs, with the parts of the proxy
|
||||||
|
// whose state the state files keep, and the metrics.
|
||||||
|
type Server struct {
|
||||||
|
*http.Server
|
||||||
|
|
||||||
|
Ledger *bans.Ledger
|
||||||
|
Limiter *ratelimit.Limiter
|
||||||
|
GeoJS *lookup.GeoJS
|
||||||
|
Metrics *metrics.Metrics
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns the server smallwebwaf runs: each request it reads passes
|
// New returns the server smallwebwaf runs: each request it reads passes
|
||||||
// through the proxy. Go's server itself refuses headers over 32 KiB, with
|
// through the proxy. Go's server itself refuses a request line and
|
||||||
// 431, closes a connection idle for 120 seconds, and applies
|
// 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
|
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
|
||||||
// applies the timeouts and size limits from then on.
|
// applies the timeouts and size limits from then on.
|
||||||
func New(params Params) *http.Server {
|
func New(params Params) *Server {
|
||||||
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
||||||
|
m := metrics.New(params.Config.MetricsTopN)
|
||||||
|
h := &handler{
|
||||||
|
config: params.Config,
|
||||||
|
requestLog: params.RequestLog,
|
||||||
|
processLog: params.ProcessLog,
|
||||||
|
errorLog: errorLog,
|
||||||
|
transport: newTransport(),
|
||||||
|
now: params.Now,
|
||||||
|
metrics: m,
|
||||||
|
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,
|
||||||
|
AttackBanDuration: params.Config.AttackBanDuration,
|
||||||
|
MaxBans: params.Config.MaxBans,
|
||||||
|
}),
|
||||||
|
geojs: lookup.New(lookup.Params{
|
||||||
|
URL: params.GeoJSURL,
|
||||||
|
Now: params.Now,
|
||||||
|
ProcessLog: params.ProcessLog,
|
||||||
|
Metrics: m,
|
||||||
|
}),
|
||||||
|
rules: params.Rules,
|
||||||
|
}
|
||||||
|
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
|
||||||
|
m.AddRules(params.Rules)
|
||||||
|
|
||||||
return &http.Server{
|
return &Server{
|
||||||
Addr: params.Config.ListenAddr,
|
Server: &http.Server{
|
||||||
Handler: &handler{
|
Addr: params.Config.ListenAddr,
|
||||||
config: params.Config,
|
Handler: h,
|
||||||
requestLog: params.RequestLog,
|
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
||||||
processLog: params.ProcessLog,
|
// Off is an IdleTimeout of 0, which Go's server replaces with
|
||||||
errorLog: errorLog,
|
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
|
||||||
transport: newTransport(),
|
IdleTimeout: params.Config.ClientIdleTimeout,
|
||||||
limiter: ratelimit.New(ratelimit.Limits{
|
// Go's server reads 4 KiB past MaxHeaderBytes before it
|
||||||
PerMinute: params.Config.RateLimitPerMinute,
|
// refuses, so the limit a client meets is the setting.
|
||||||
PerHour: params.Config.RateLimitPerHour,
|
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
||||||
PerDay: params.Config.RateLimitPerDay,
|
ErrorLog: errorLog,
|
||||||
}),
|
|
||||||
geojs: lookup.New(lookup.Params{
|
|
||||||
URL: params.GeoJSURL,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: params.ProcessLog,
|
|
||||||
}),
|
|
||||||
},
|
},
|
||||||
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
Ledger: h.ledger,
|
||||||
IdleTimeout: clientIdleTimeout,
|
Limiter: h.limiter,
|
||||||
MaxHeaderBytes: requestHeaderMaxBytes,
|
GeoJS: h.geojs,
|
||||||
ErrorLog: errorLog,
|
Metrics: m,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,8 +136,12 @@ type handler struct {
|
|||||||
processLog *slog.Logger
|
processLog *slog.Logger
|
||||||
errorLog *log.Logger
|
errorLog *log.Logger
|
||||||
transport http.RoundTripper
|
transport http.RoundTripper
|
||||||
|
now func() time.Time
|
||||||
|
metrics *metrics.Metrics
|
||||||
limiter *ratelimit.Limiter
|
limiter *ratelimit.Limiter
|
||||||
|
ledger *bans.Ledger
|
||||||
geojs *lookup.GeoJS
|
geojs *lookup.GeoJS
|
||||||
|
rules *rules.Files
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTransport returns what carries requests to the app. It never goes
|
// newTransport returns what carries requests to the app. It never goes
|
||||||
@@ -111,7 +158,8 @@ func newTransport() *http.Transport {
|
|||||||
|
|
||||||
// ServeHTTP handles one request: it works out the client, runs the
|
// ServeHTTP handles one request: it works out the client, runs the
|
||||||
// checks, passes the request to the app and the answer back within the
|
// checks, passes the request to the app and the answer back within the
|
||||||
// limits, and writes the request's log line.
|
// limits, or answers it itself if it is for smallwebwaf, and writes the
|
||||||
|
// request's log line.
|
||||||
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
rq := h.newRequest(w, r)
|
rq := h.newRequest(w, r)
|
||||||
defer rq.finish()
|
defer rq.finish()
|
||||||
@@ -120,17 +168,33 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
// a health checker is never refused. It does not ask the app.
|
// a health checker is never refused. It does not ask the app.
|
||||||
if r.Method == http.MethodGet && r.URL.Path == HealthPath {
|
if r.Method == http.MethodGet && r.URL.Path == HealthPath {
|
||||||
rq.line.Action = requestlog.ActionAdmin
|
rq.line.Action = requestlog.ActionAdmin
|
||||||
|
// Set here rather than left to Go's server, which would set it only
|
||||||
|
// after the log line has taken the response's headers.
|
||||||
|
rq.out.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||||
_, _ = io.WriteString(rq.out, "ok\n")
|
_, _ = io.WriteString(rq.out, "ok\n")
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Once the request has ended, before its log line is written.
|
||||||
|
defer rq.addToHistory()
|
||||||
|
|
||||||
refused := rq.check(r.Context())
|
refused := rq.check(r.Context())
|
||||||
|
rq.checked = time.Now()
|
||||||
|
|
||||||
if refused != nil {
|
if refused != nil {
|
||||||
rq.answer(*refused)
|
rq.answer(*refused)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A request for smallwebwaf itself is answered where another would be
|
||||||
|
// passed to the app, so that it goes through every check first.
|
||||||
|
if strings.HasPrefix(r.URL.Path, adminPrefix) {
|
||||||
|
rq.answerAdmin()
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
rq.forward(r.Context())
|
rq.forward(r.Context())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -35,6 +36,10 @@ const (
|
|||||||
// localhost is where every test server listens, and so the address
|
// localhost is where every test server listens, and so the address
|
||||||
// smallwebwaf sees each test's requests come from.
|
// smallwebwaf sees each test's requests come from.
|
||||||
localhost = "127.0.0.1"
|
localhost = "127.0.0.1"
|
||||||
|
// requestType is the type that marks a request log line.
|
||||||
|
requestType = "request"
|
||||||
|
// protocol is the protocol of every test's requests.
|
||||||
|
protocol = "HTTP/1.1"
|
||||||
)
|
)
|
||||||
|
|
||||||
// shortTimeoutSetting is shortTimeout as a setting's value.
|
// shortTimeoutSetting is shortTimeout as a setting's value.
|
||||||
@@ -45,9 +50,12 @@ var shortTimeoutSetting = shortTimeout.String()
|
|||||||
// The settings the tests set.
|
// The settings the tests set.
|
||||||
const (
|
const (
|
||||||
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
||||||
|
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
|
||||||
|
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
|
||||||
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
||||||
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
||||||
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
||||||
|
mode = "SWWAF_MODE"
|
||||||
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
|
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
|
||||||
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
||||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||||
@@ -55,8 +63,20 @@ const (
|
|||||||
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
|
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
|
||||||
denyNets = "SWWAF_DENY_NETS"
|
denyNets = "SWWAF_DENY_NETS"
|
||||||
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
||||||
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
|
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
||||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||||
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_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"
|
||||||
|
instanceName = "SWWAF_INSTANCE_NAME"
|
||||||
|
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
|
||||||
|
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
|
||||||
|
rulesDir = "SWWAF_RULES_DIR"
|
||||||
)
|
)
|
||||||
|
|
||||||
// output collects what smallwebwaf writes on stdout.
|
// output collects what smallwebwaf writes on stdout.
|
||||||
@@ -73,6 +93,14 @@ func (o *output) Write(p []byte) (int, error) {
|
|||||||
return o.buf.Write(p)
|
return o.buf.Write(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// text returns everything written so far.
|
||||||
|
func (o *output) text() string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
return o.buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
// lines returns every line written so far, decoded.
|
// lines returns every line written so far, decoded.
|
||||||
func (o *output) lines(t *testing.T) []map[string]any {
|
func (o *output) lines(t *testing.T) []map[string]any {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
@@ -112,7 +140,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
|
|||||||
var found []logLine
|
var found []logLine
|
||||||
|
|
||||||
for _, fields := range o.lines(t) {
|
for _, fields := range o.lines(t) {
|
||||||
if fields["type"] == "request" {
|
if fields["type"] == requestType {
|
||||||
found = append(found, decodeLine(t, fields))
|
found = append(found, decodeLine(t, fields))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -181,7 +209,21 @@ func startProxyWithGeoJS(
|
|||||||
) (string, *output) {
|
) (string, *output) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
|
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. Unless
|
||||||
|
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
|
||||||
|
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, rulesDir: t.TempDir()}
|
||||||
maps.Copy(settings, env)
|
maps.Copy(settings, env)
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
|
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
|
||||||
@@ -194,11 +236,22 @@ func startProxyWithGeoJS(
|
|||||||
}
|
}
|
||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
|
processLog := requestlog.NewProcessLogger(out)
|
||||||
|
|
||||||
|
ruleFiles, err := rules.Load(rules.Params{
|
||||||
|
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("rule files: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
server := proxy.New(proxy.Params{
|
||||||
Config: cfg,
|
Config: cfg,
|
||||||
RequestLog: out,
|
RequestLog: out,
|
||||||
ProcessLog: requestlog.NewProcessLogger(out),
|
ProcessLog: processLog,
|
||||||
GeoJSURL: geojsURL,
|
GeoJSURL: geojsURL,
|
||||||
|
Now: now,
|
||||||
|
Rules: ruleFiles,
|
||||||
})
|
})
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
@@ -214,7 +267,7 @@ func startProxyWithGeoJS(
|
|||||||
_ = server.Close()
|
_ = 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,
|
// newClient returns an HTTP client that sends requests as they are made,
|
||||||
|
|||||||
@@ -5,10 +5,16 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"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()
|
t.Parallel()
|
||||||
|
|
||||||
var calls atomic.Int32
|
var calls atomic.Int32
|
||||||
@@ -24,19 +30,19 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
|
|||||||
const otherClient = "203.0.113.10"
|
const otherClient = "203.0.113.10"
|
||||||
|
|
||||||
// With a limit of one request a minute, a client's second request is
|
// 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
|
// refused, with 403 by default. A client is one IPv4 address, or one
|
||||||
// address in IPv6 form is that IPv4 address.
|
// IPv6 /64; an IPv4 address in IPv6 form is that IPv4 address.
|
||||||
requests := []struct {
|
requests := []struct {
|
||||||
client string // as X-Forwarded-For names it
|
client string // as X-Forwarded-For names it
|
||||||
logged string // as the log line's client_ip names it
|
logged string // as the log line's client_ip names it
|
||||||
want int
|
want int
|
||||||
}{
|
}{
|
||||||
{client, client, http.StatusOK},
|
{client, client, http.StatusOK},
|
||||||
{client, client, http.StatusTooManyRequests},
|
{client, client, http.StatusForbidden},
|
||||||
{otherClient, otherClient, http.StatusOK},
|
{otherClient, otherClient, http.StatusOK},
|
||||||
{"::ffff:" + otherClient, otherClient, http.StatusTooManyRequests},
|
{"::ffff:" + otherClient, otherClient, http.StatusForbidden},
|
||||||
{"2001:db8::1", "2001:db8::1", http.StatusOK},
|
{"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},
|
{"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -53,9 +59,9 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
|
|||||||
if sent.want == http.StatusOK {
|
if sent.want == http.StatusOK {
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
||||||
} else {
|
} 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)
|
t.Errorf("log line has limit_hit %q, want minute", line.LimitHit)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -65,3 +71,89 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
|
|||||||
t.Errorf("the app was called %d times, want 4", calls.Load())
|
t.Errorf("the app was called %d times, want 4", calls.Load())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
|
||||||
|
s, _, server := startWithClock(t, "", map[string]string{
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
rateLimitExemptPaths: "/assets/,/favicon.ico",
|
||||||
|
denyNets: denied,
|
||||||
|
deniedCountries: "kp",
|
||||||
|
})
|
||||||
|
|
||||||
|
// The answers are kept before the requests, so that none waits for
|
||||||
|
// GeoJS.
|
||||||
|
server.GeoJS.Load([]lookup.Answer{
|
||||||
|
keptAnswer(client, "DE"), keptAnswer(fromKP, "KP"),
|
||||||
|
})
|
||||||
|
|
||||||
|
// With a limit of one request a minute, the requests for paths under a
|
||||||
|
// prefix are not counted, so client's first request for / is within
|
||||||
|
// the limit; and once client has reached it, they are not refused.
|
||||||
|
s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
|
||||||
|
if line.LimitHit != "" || line.Counts != (ratelimit.Counts{}) {
|
||||||
|
t.Errorf("log line has limit_hit %q and counts %+v, want neither",
|
||||||
|
line.LimitHit, line.Counts)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A path outside every prefix is counted: /assets is not under
|
||||||
|
// /assets/, and breaks the limit.
|
||||||
|
s.request(client, "/assets", http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
|
// A ban, SWWAF_DENY_NETS and the country lists still refuse a path
|
||||||
|
// under a prefix.
|
||||||
|
s.request(client, "/assets/app.js", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
s.request(denied, "/assets/app.js", http.StatusForbidden, requestlog.ActionDenied)
|
||||||
|
s.request(fromKP, "/assets/app.js",
|
||||||
|
http.StatusForbidden, requestlog.ActionCountryDenied)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRateLimitCountsPathsThatAreNotExempt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, sent := range []string{
|
||||||
|
// A prefix matches only at the start of the path.
|
||||||
|
"/static/assets/app.js",
|
||||||
|
// A prefix matches the path as sent: a router that matches the
|
||||||
|
// path as received does not take /%61ssets/x for a path under
|
||||||
|
// /assets/.
|
||||||
|
"/%61ssets/x",
|
||||||
|
// .. once percent-decoded: an app may act on these as /login, the
|
||||||
|
// last as a path under /sneak/app/ or as /assets/x.
|
||||||
|
"/assets/../login",
|
||||||
|
"/assets/%2e%2e/login",
|
||||||
|
"/assets/..%2Flogin",
|
||||||
|
"/assets/..;/login",
|
||||||
|
"/sneak/app/src/branch/main/..%2F..%2F..%2F..%2F..%2F..%2Fassets/x",
|
||||||
|
// Not under /assets/ as sent: Go's router takes /assets%2Fx for one
|
||||||
|
// path segment, not a path under /assets/.
|
||||||
|
"/assets%2Fx",
|
||||||
|
"/assets%2fx",
|
||||||
|
// Under /assets/ as sent, but holding an encoded slash, in either
|
||||||
|
// case, or a backslash: never exempt, whatever the prefix.
|
||||||
|
"/assets/x%2Fy",
|
||||||
|
"/assets/x%2fy",
|
||||||
|
`/assets/x\y`,
|
||||||
|
} {
|
||||||
|
t.Run(sent, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, _ := startWithClock(t, "", map[string]string{
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
rateLimitExemptPaths: "/assets/",
|
||||||
|
})
|
||||||
|
|
||||||
|
// Counted, the second request breaks the limit of one request
|
||||||
|
// a minute.
|
||||||
|
s.request(client, sent, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.request(client, sent, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+269
-69
@@ -7,11 +7,15 @@ import (
|
|||||||
"net/http/httptrace"
|
"net/http/httptrace"
|
||||||
"net/http/httputil"
|
"net/http/httputil"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -21,10 +25,13 @@ const flushAfterEachWrite time.Duration = -1
|
|||||||
|
|
||||||
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
|
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
|
||||||
// the status the client is answered if the response has not started yet,
|
// the status the client is answered if the response has not started yet,
|
||||||
// and the action the log line names.
|
// 0 to close the connection without an answer, the action the log line
|
||||||
|
// names, and the setting whose size or time limit the request passed, if
|
||||||
|
// that is why.
|
||||||
type refusal struct {
|
type refusal struct {
|
||||||
status int
|
status int
|
||||||
action string
|
action string
|
||||||
|
limit string
|
||||||
}
|
}
|
||||||
|
|
||||||
// request is one request on its way through smallwebwaf, from the moment
|
// request is one request on its way through smallwebwaf, from the moment
|
||||||
@@ -42,7 +49,9 @@ type request struct {
|
|||||||
peer netip.Addr
|
peer netip.Addr
|
||||||
peerTrusted bool
|
peerTrusted bool
|
||||||
start time.Time
|
start time.Time
|
||||||
// upstreamStart is when the request was handed to the app.
|
// checked is when the checks were done, and upstreamStart when the
|
||||||
|
// request was handed to the app.
|
||||||
|
checked time.Time
|
||||||
upstreamStart time.Time
|
upstreamStart time.Time
|
||||||
// cancel ends the request to the app.
|
// cancel ends the request to the app.
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
@@ -52,24 +61,34 @@ type request struct {
|
|||||||
complete bool
|
complete bool
|
||||||
|
|
||||||
// mu guards what follows. The timeouts run on goroutines of their
|
// mu guards what follows. The timeouts run on goroutines of their
|
||||||
// own, and the transport starts and stops them from its own; once
|
// own, and the transport starts and stops them, and notes the times
|
||||||
// timersStopped is set, none of them acts any more.
|
// below, from its own; once timersStopped is set, none of the timeouts
|
||||||
|
// acts any more.
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
timersStopped bool
|
timersStopped bool
|
||||||
clientRequestTimer *time.Timer
|
clientRequestTimer *time.Timer
|
||||||
upstreamRequestTimer *time.Timer
|
upstreamRequestTimer *time.Timer
|
||||||
upstreamResponseTimer *time.Timer
|
upstreamResponseTimer *time.Timer
|
||||||
// requestSent is when the app had been sent the whole request.
|
// connected is when there was a connection to the app, requestSent
|
||||||
requestSent time.Time
|
// when the app had been sent the whole request, and answerStarted
|
||||||
|
// when the first byte of its answer arrived.
|
||||||
|
connected time.Time
|
||||||
|
requestSent time.Time
|
||||||
|
answerStarted time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
// newRequest starts handling r: it notes the time and works out the
|
// newRequest starts handling r: it notes the time, counts the request as
|
||||||
// client.
|
// under way, works out the client, and starts the log line with what is
|
||||||
|
// known of the request.
|
||||||
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
||||||
|
h.metrics.RequestStarted()
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
peer := peerAddress(r)
|
peer := peerAddress(r)
|
||||||
trusted := h.config.TrustedProxies
|
trusted := h.config.TrustedProxies
|
||||||
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted)
|
peerTrusted := isInside(peer, trusted)
|
||||||
|
forwardedFor := r.Header.Values("X-Forwarded-For")
|
||||||
|
client := clientAddress(peer, forwardedFor, trusted)
|
||||||
|
|
||||||
rq := &request{
|
rq := &request{
|
||||||
h: h,
|
h: h,
|
||||||
@@ -78,22 +97,37 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
|||||||
out: &responseWriter{ResponseWriter: w},
|
out: &responseWriter{ResponseWriter: w},
|
||||||
client: client,
|
client: client,
|
||||||
peer: peer,
|
peer: peer,
|
||||||
peerTrusted: isInside(peer, trusted),
|
peerTrusted: peerTrusted,
|
||||||
start: start,
|
start: start,
|
||||||
line: requestlog.Line{
|
line: requestlog.Line{
|
||||||
Time: requestlog.FormatTime(start),
|
Time: requestlog.FormatTime(start),
|
||||||
ClientIP: client.String(),
|
Instance: h.config.InstanceName,
|
||||||
PeerIP: peer.String(),
|
ClientIP: client.String(),
|
||||||
Method: r.Method,
|
Method: r.Method,
|
||||||
Host: r.Host,
|
Scheme: scheme(r, peerTrusted),
|
||||||
Path: r.URL.EscapedPath(),
|
Host: r.Host,
|
||||||
Query: r.URL.RawQuery,
|
Path: r.URL.EscapedPath(),
|
||||||
Protocol: r.Proto,
|
Query: r.URL.RawQuery,
|
||||||
Referer: r.Referer(),
|
Protocol: r.Proto,
|
||||||
UserAgent: r.UserAgent(),
|
Referer: r.Referer(),
|
||||||
Action: requestlog.ActionForward,
|
UserAgent: r.UserAgent(),
|
||||||
|
RequestID: requestID(r, peerTrusted),
|
||||||
|
PeerIP: peer.String(),
|
||||||
|
ForwardedFor: strings.Join(forwardedFor, ", "),
|
||||||
|
ClientGroup: clientGroup(client).String(),
|
||||||
|
ContentType: r.Header.Get("Content-Type"),
|
||||||
|
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
|
||||||
|
HasAuthorization: len(r.Header.Values("Authorization")) > 0,
|
||||||
|
HasCookie: len(r.Header.Values("Cookie")) > 0,
|
||||||
|
Action: requestlog.ActionForward,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A length of -1 is a body whose length was not announced.
|
||||||
|
if r.ContentLength > 0 {
|
||||||
|
rq.line.ContentLength = r.ContentLength
|
||||||
|
}
|
||||||
|
|
||||||
if r.Body != http.NoBody {
|
if r.Body != http.NoBody {
|
||||||
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
|
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
|
||||||
}
|
}
|
||||||
@@ -101,56 +135,124 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
|||||||
return rq
|
return rq
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// requestHeaders returns the headers of r that names lists, by name in
|
||||||
|
// lower case, each with its values joined by ", ". Authorization, Cookie
|
||||||
|
// and Set-Cookie are never among them, whatever names says.
|
||||||
|
func requestHeaders(r *http.Request, names []string) map[string]string {
|
||||||
|
headers := map[string]string{}
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
|
switch name {
|
||||||
|
case "authorization", "cookie", "set-cookie":
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
values := r.Header.Values(name)
|
||||||
|
if len(values) > 0 {
|
||||||
|
headers[name] = strings.Join(values, ", ")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return headers
|
||||||
|
}
|
||||||
|
|
||||||
// check is the one place where a request can be refused once its client
|
// check is the one place where a request can be refused once its client
|
||||||
// is known, before its body is read or anything reaches the app. It
|
// is known, before its body is read or anything reaches the app. It
|
||||||
// returns nil to let the request through. A client in SWWAF_ALLOW_NETS
|
// returns nil to let the request through. The checks of checkClient come
|
||||||
// skips every check but the size limit. For any other client,
|
// first, answered with SWWAF_BAN_RESPONSE, or 403 for a block rule, and
|
||||||
// SWWAF_DENY_NETS comes first, so that a client it refuses is not looked
|
// then the size limit, so that a request the rate limits count is counted
|
||||||
// up, and then the country lists; a request either refuses is not counted
|
// even when it is refused for its size. In observe mode a request
|
||||||
// for the rate limits. Then come the rate limits, unless the client is in
|
// checkClient refuses goes on to the size limit like any other. ctx is
|
||||||
// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted,
|
// the request's own context.
|
||||||
// one refused for its size too. ctx is the request's own context.
|
|
||||||
func (rq *request) check(ctx context.Context) *refusal {
|
func (rq *request) check(ctx context.Context) *refusal {
|
||||||
cfg := rq.h.config
|
action := rq.checkClient(ctx)
|
||||||
allowed := isInside(rq.client, cfg.AllowNets)
|
|
||||||
|
|
||||||
if !allowed && isInside(rq.client, cfg.DenyNets) {
|
switch {
|
||||||
return &refusal{
|
case action == "":
|
||||||
status: http.StatusForbidden,
|
case rq.h.config.Observe:
|
||||||
action: requestlog.ActionDenied,
|
// The log line names what enforce mode would have done.
|
||||||
}
|
rq.line.WouldAction = action
|
||||||
|
case action == requestlog.ActionRuleBlocked:
|
||||||
|
return &refusal{status: http.StatusForbidden, action: action}
|
||||||
|
default:
|
||||||
|
return rq.banResponse(action)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !allowed && rq.countryDenied(ctx) {
|
maxBytes := rq.h.config.RequestMaxBytes
|
||||||
return &refusal{
|
|
||||||
status: http.StatusForbidden,
|
|
||||||
action: 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,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
maxBytes := cfg.RequestMaxBytes
|
|
||||||
if maxBytes > 0 && rq.in.ContentLength > maxBytes {
|
if maxBytes > 0 && rq.in.ContentLength > maxBytes {
|
||||||
return &refusal{
|
return &refusal{
|
||||||
status: http.StatusRequestEntityTooLarge,
|
status: http.StatusRequestEntityTooLarge,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
|
limit: "SWWAF_REQUEST_MAX_BYTES",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// checkClient runs the checks on the request's client, and returns the
|
||||||
|
// action of the first that refuses the request, or "" when none does. A
|
||||||
|
// client in SWWAF_ALLOW_NETS skips them. For any other client,
|
||||||
|
// 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 or the request's path is exempt under
|
||||||
|
// SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted,
|
||||||
|
// and last the rule files. ctx is the request's own context.
|
||||||
|
func (rq *request) checkClient(ctx context.Context) string {
|
||||||
|
cfg := rq.h.config
|
||||||
|
if isInside(rq.client, cfg.AllowNets) {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
now := rq.h.now()
|
||||||
|
|
||||||
|
if isInside(rq.client, cfg.DenyNets) {
|
||||||
|
return requestlog.ActionDenied
|
||||||
|
}
|
||||||
|
|
||||||
|
if rq.banned(now) {
|
||||||
|
return requestlog.ActionBanned
|
||||||
|
}
|
||||||
|
|
||||||
|
if rq.countryDenied(ctx) {
|
||||||
|
return requestlog.ActionCountryDenied
|
||||||
|
}
|
||||||
|
|
||||||
|
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
|
||||||
|
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
||||||
|
if !exempt && rq.limitBroken(now) {
|
||||||
|
return requestlog.ActionRateLimited
|
||||||
|
}
|
||||||
|
|
||||||
|
return rq.checkRules(now)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pathExempt reports whether the rate limits leave out a request for u
|
||||||
|
// because of SWWAF_RATE_LIMIT_EXEMPT_PATHS: whether its path as sent, the
|
||||||
|
// path the app receives, not percent-decoded, starts with one of
|
||||||
|
// prefixes, so that /%61ssets/x is not under /assets/ for an app whose
|
||||||
|
// router matches the path as received. A request whose decoded path
|
||||||
|
// contains .. anywhere or a backslash, or whose path as sent holds an
|
||||||
|
// encoded slash (%2F or %2f), never is, since an app may act on it as a
|
||||||
|
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
|
||||||
|
// as one path segment, as Go's router does.
|
||||||
|
func pathExempt(u *url.URL, prefixes []string) bool {
|
||||||
|
decoded := u.Path
|
||||||
|
// EscapedPath is the path as the app receives it, not decoded.
|
||||||
|
sent := u.EscapedPath()
|
||||||
|
|
||||||
|
if strings.Contains(decoded, "..") || strings.Contains(decoded, `\`) ||
|
||||||
|
strings.Contains(strings.ToLower(sent), "%2f") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return slices.ContainsFunc(prefixes, func(prefix string) bool {
|
||||||
|
return strings.HasPrefix(sent, prefix)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// forward passes the request to the app and the app's answer back. ctx
|
// forward passes the request to the app and the app's answer back. ctx
|
||||||
// is the request's own context.
|
// is the request's own context.
|
||||||
func (rq *request) forward(ctx context.Context) {
|
func (rq *request) forward(ctx context.Context) {
|
||||||
@@ -159,7 +261,9 @@ func (rq *request) forward(ctx context.Context) {
|
|||||||
|
|
||||||
rq.cancel = cancel
|
rq.cancel = cancel
|
||||||
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
|
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
|
||||||
WroteRequest: rq.wroteRequest,
|
GotConn: rq.gotConn,
|
||||||
|
WroteRequest: rq.wroteRequest,
|
||||||
|
GotFirstResponseByte: rq.gotFirstResponseByte,
|
||||||
})
|
})
|
||||||
|
|
||||||
out := rq.in.WithContext(ctx)
|
out := rq.in.WithContext(ctx)
|
||||||
@@ -182,7 +286,8 @@ func (rq *request) forward(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// rewrite makes the request the app receives: the client's request,
|
// rewrite makes the request the app receives: the client's request,
|
||||||
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set.
|
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
|
||||||
|
// the request's id set.
|
||||||
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
||||||
upstream := rq.h.config.UpstreamURL
|
upstream := rq.h.config.UpstreamURL
|
||||||
pr.Out.URL.Scheme = upstream.Scheme
|
pr.Out.URL.Scheme = upstream.Scheme
|
||||||
@@ -191,6 +296,7 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
|||||||
// the query as the client sent it.
|
// the query as the client sent it.
|
||||||
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
||||||
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
||||||
|
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
|
||||||
}
|
}
|
||||||
|
|
||||||
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
||||||
@@ -204,13 +310,18 @@ func (rq *request) modifyResponse(res *http.Response) error {
|
|||||||
// connection it takes over, not through rq.out.
|
// connection it takes over, not through rq.out.
|
||||||
rq.stopTimers()
|
rq.stopTimers()
|
||||||
rq.out.status = res.StatusCode
|
rq.out.status = res.StatusCode
|
||||||
|
rq.line.Websocket = true
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
maxBytes := rq.h.config.ResponseMaxBytes
|
maxBytes := rq.h.config.ResponseMaxBytes
|
||||||
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
|
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
|
||||||
rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge})
|
rq.refuse(refusal{
|
||||||
|
status: http.StatusBadGateway,
|
||||||
|
action: requestlog.ActionTooLarge,
|
||||||
|
limit: "SWWAF_RESPONSE_MAX_BYTES",
|
||||||
|
})
|
||||||
|
|
||||||
return errResponseTooLarge
|
return errResponseTooLarge
|
||||||
}
|
}
|
||||||
@@ -250,6 +361,13 @@ func (rq *request) answer(r refusal) {
|
|||||||
return // too late to answer: the connection can only be cut
|
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
|
// 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
|
// sending until its time is up, so that Go's server can read the
|
||||||
// rest of the body and end the request cleanly.
|
// rest of the body and end the request cleanly.
|
||||||
@@ -275,7 +393,8 @@ func (rq *request) refuse(r refusal) {
|
|||||||
rq.cancel()
|
rq.cancel()
|
||||||
}
|
}
|
||||||
|
|
||||||
// finish ends the request's timeouts and writes its log line.
|
// finish ends the request's timeouts, counts it in the metrics and writes
|
||||||
|
// its log line.
|
||||||
func (rq *request) finish() {
|
func (rq *request) finish() {
|
||||||
rq.stopTimers()
|
rq.stopTimers()
|
||||||
|
|
||||||
@@ -287,35 +406,90 @@ func (rq *request) finish() {
|
|||||||
line := &rq.line
|
line := &rq.line
|
||||||
line.Status = rq.out.status
|
line.Status = rq.out.status
|
||||||
line.ResponseBytes = rq.out.bytes
|
line.ResponseBytes = rq.out.bytes
|
||||||
|
header := rq.out.Header()
|
||||||
|
line.ResponseContentType = header.Get("Content-Type")
|
||||||
|
line.CacheControl = header.Get("Cache-Control")
|
||||||
|
line.Location = header.Get("Location")
|
||||||
|
|
||||||
if rq.body != nil {
|
if rq.body != nil {
|
||||||
line.RequestBytes = rq.body.bytes.Load()
|
line.RequestBytes = rq.body.bytes.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// limit is the setting whose size or time limit the request passed.
|
||||||
|
var limit string
|
||||||
|
|
||||||
switch {
|
switch {
|
||||||
case refused != nil:
|
case refused != nil:
|
||||||
line.Action = refused.action
|
line.Action = refused.action
|
||||||
|
limit = refused.limit
|
||||||
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
|
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
|
||||||
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
|
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
|
||||||
// take the response.
|
// take the response.
|
||||||
line.Action = requestlog.ActionTimedOut
|
line.Action = requestlog.ActionTimedOut
|
||||||
|
limit = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
||||||
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
|
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
|
||||||
line.Aborted = true
|
line.Aborted = true
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start))
|
duration := now.Sub(rq.start)
|
||||||
|
line.DurationTotal = requestlog.Milliseconds(duration)
|
||||||
|
line.DurationChecks = timing(rq.start, rq.checked)
|
||||||
|
|
||||||
|
var upstreamDuration time.Duration
|
||||||
|
|
||||||
if !rq.upstreamStart.IsZero() {
|
if !rq.upstreamStart.IsZero() {
|
||||||
line.DurationUpstreamTotal = requestlog.Milliseconds(now.Sub(rq.upstreamStart))
|
upstreamDuration = now.Sub(rq.upstreamStart)
|
||||||
|
line.DurationUpstreamTotal = new(requestlog.Milliseconds(upstreamDuration))
|
||||||
|
|
||||||
|
rq.mu.Lock()
|
||||||
|
line.DurationUpstreamConnect = timing(rq.upstreamStart, rq.connected)
|
||||||
|
line.DurationUpstreamFirstByte = timing(rq.upstreamStart, rq.answerStarted)
|
||||||
|
rq.mu.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Counted before the log line is written, so that the metrics count
|
||||||
|
// every request whose line is out.
|
||||||
|
rq.h.metrics.RequestEnded(line, limit, duration, upstreamDuration)
|
||||||
|
|
||||||
err := requestlog.Write(rq.h.requestLog, line)
|
err := requestlog.Write(rq.h.requestLog, line)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
rq.h.processLog.Error("writing the request log failed", "error", err.Error())
|
rq.h.processLog.Error("writing the request log failed", "error", err.Error())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// timing is the time from start to end in milliseconds, for one of the
|
||||||
|
// log line's timings, or nil when end is zero: what it times never
|
||||||
|
// happened.
|
||||||
|
func timing(start, end time.Time) *float64 {
|
||||||
|
if end.IsZero() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return new(requestlog.Milliseconds(end.Sub(start)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// addToHistory adds the request, which has ended, to its client's
|
||||||
|
// history.
|
||||||
|
func (rq *request) addToHistory() {
|
||||||
|
var requestBytes int64
|
||||||
|
if rq.body != nil {
|
||||||
|
requestBytes = rq.body.bytes.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
forwarded := !rq.upstreamStart.IsZero()
|
||||||
|
|
||||||
|
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
|
||||||
|
Country: rq.line.Country,
|
||||||
|
Forwarded: forwarded,
|
||||||
|
Refused: !forwarded && rq.refused.Load() != nil,
|
||||||
|
Status: rq.out.status,
|
||||||
|
RequestBytes: requestBytes,
|
||||||
|
ResponseBytes: rq.out.bytes,
|
||||||
|
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// clientRequestDeadline is when the client must have sent its whole
|
// clientRequestDeadline is when the client must have sent its whole
|
||||||
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
||||||
func (rq *request) clientRequestDeadline() time.Time {
|
func (rq *request) clientRequestDeadline() time.Time {
|
||||||
@@ -348,21 +522,26 @@ func (rq *request) startRequestTimers() {
|
|||||||
|
|
||||||
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
|
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
|
||||||
rq.clientRequestTimer = time.AfterFunc(
|
rq.clientRequestTimer = time.AfterFunc(
|
||||||
time.Until(rq.clientRequestDeadline()), rq.requestTimedOut)
|
time.Until(rq.clientRequestDeadline()), func() {
|
||||||
|
rq.requestTimedOut("SWWAF_CLIENT_REQUEST_TIMEOUT")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
timeout := rq.h.config.UpstreamRequestTimeout
|
timeout := rq.h.config.UpstreamRequestTimeout
|
||||||
if timeout > 0 {
|
if timeout > 0 {
|
||||||
rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut)
|
rq.upstreamRequestTimer = time.AfterFunc(timeout, func() {
|
||||||
|
rq.requestTimedOut("SWWAF_UPSTREAM_REQUEST_TIMEOUT")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// requestTimedOut is called when a request timeout runs out while the
|
// requestTimedOut is called when limit, SWWAF_CLIENT_REQUEST_TIMEOUT or
|
||||||
// request is still on its way to the app. The answer names the side
|
// SWWAF_UPSTREAM_REQUEST_TIMEOUT, runs out while the request is still on
|
||||||
// smallwebwaf was waiting on at that moment: 408 when it was waiting for
|
// its way to the app. The answer names the side smallwebwaf was waiting
|
||||||
// the client to send more of its body, 504 when it was waiting for the
|
// on at that moment: 408 when it was waiting for the client to send more
|
||||||
// app to be reached or to take what it had.
|
// of its body, 504 when it was waiting for the app to be reached or to
|
||||||
func (rq *request) requestTimedOut() {
|
// take what it had.
|
||||||
|
func (rq *request) requestTimedOut(limit string) {
|
||||||
rq.mu.Lock()
|
rq.mu.Lock()
|
||||||
defer rq.mu.Unlock()
|
defer rq.mu.Unlock()
|
||||||
|
|
||||||
@@ -374,6 +553,7 @@ func (rq *request) requestTimedOut() {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusGatewayTimeout,
|
status: http.StatusGatewayTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
|
limit: limit,
|
||||||
})
|
})
|
||||||
|
|
||||||
return
|
return
|
||||||
@@ -382,6 +562,7 @@ func (rq *request) requestTimedOut() {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusRequestTimeout,
|
status: http.StatusRequestTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
|
limit: limit,
|
||||||
})
|
})
|
||||||
// The transport gives up on the app only once its Read of the
|
// The transport gives up on the app only once its Read of the
|
||||||
// client's body returns, so that Read is ended now. The lock keeps
|
// client's body returns, so that Read is ended now. The lock keeps
|
||||||
@@ -397,6 +578,24 @@ func (rq *request) bodyReceived() {
|
|||||||
stopTimer(rq.clientRequestTimer)
|
stopTimer(rq.clientRequestTimer)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// gotConn is called once there is a connection to the app, a new one or
|
||||||
|
// one kept open from an earlier request.
|
||||||
|
func (rq *request) gotConn(httptrace.GotConnInfo) {
|
||||||
|
rq.mu.Lock()
|
||||||
|
defer rq.mu.Unlock()
|
||||||
|
|
||||||
|
rq.connected = time.Now()
|
||||||
|
}
|
||||||
|
|
||||||
|
// gotFirstResponseByte is called once the first byte of the app's answer
|
||||||
|
// has arrived.
|
||||||
|
func (rq *request) gotFirstResponseByte() {
|
||||||
|
rq.mu.Lock()
|
||||||
|
defer rq.mu.Unlock()
|
||||||
|
|
||||||
|
rq.answerStarted = time.Now()
|
||||||
|
}
|
||||||
|
|
||||||
// wroteRequest is called once the app has been sent the whole request:
|
// wroteRequest is called once the app has been sent the whole request:
|
||||||
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
|
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
|
||||||
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
|
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
|
||||||
@@ -431,6 +630,7 @@ func (rq *request) responseTimedOut() {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusGatewayTimeout,
|
status: http.StatusGatewayTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
|
limit: "SWWAF_UPSTREAM_RESPONSE_TIMEOUT",
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,368 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"maps"
|
||||||
|
"math"
|
||||||
|
"net/http"
|
||||||
|
"reflect"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// requestIDHeader carries the request's id.
|
||||||
|
requestIDHeader = "X-Request-ID"
|
||||||
|
// instance is the SWWAF_INSTANCE_NAME a test sets.
|
||||||
|
instance = "fsn1app1/gitea"
|
||||||
|
// ipv6Client is a client on IPv6, and ipv6Group the netblock the rate
|
||||||
|
// limits count it as.
|
||||||
|
ipv6Client = "2001:db8::7"
|
||||||
|
ipv6Group = "2001:db8::/64"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLogLineHasEachFieldWhereItApplies(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
received := make(chan string, 2) // the request ids the app received
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
received <- r.Header.Get(requestIDHeader)
|
||||||
|
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
|
||||||
|
if r.URL.Path != "/full" {
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "text/html")
|
||||||
|
w.Header().Set("Cache-Control", "no-store")
|
||||||
|
w.Header().Set("Location", "/elsewhere")
|
||||||
|
w.WriteHeader(http.StatusFound)
|
||||||
|
_, _ = io.WriteString(w, "moved")
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
rateLimitExemptNets: localhost,
|
||||||
|
instanceName: instance,
|
||||||
|
logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE",
|
||||||
|
})
|
||||||
|
|
||||||
|
// This request comes from ipv6Client through a trusted proxy, with a
|
||||||
|
// body and each header the log line looks at, and is answered with a
|
||||||
|
// redirect.
|
||||||
|
conn := dial(t, addr)
|
||||||
|
send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+
|
||||||
|
forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+
|
||||||
|
forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+
|
||||||
|
"Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+
|
||||||
|
"Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+
|
||||||
|
"Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+
|
||||||
|
"Set-Cookie: secret-set-cookie\r\n\r\na=b")
|
||||||
|
wantStatus(t, readResponse(t, conn), http.StatusFound)
|
||||||
|
|
||||||
|
// A request's log line can come after its answer: each is waited for
|
||||||
|
// before the next request, so that the lines are in order.
|
||||||
|
full := out.requestLines(t, 1)[0]
|
||||||
|
|
||||||
|
// This one comes from 127.0.0.1, which the rate limits do not count,
|
||||||
|
// with a body of 4 bytes whose length it does not announce, so that its
|
||||||
|
// request_bytes is not its content_length, and no header the log line
|
||||||
|
// looks at, and is answered with 204 and no header.
|
||||||
|
conn = dial(t, addr)
|
||||||
|
send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+
|
||||||
|
"Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n")
|
||||||
|
wantStatus(t, readResponse(t, conn), http.StatusNoContent)
|
||||||
|
|
||||||
|
bare := out.requestLines(t, 2)[1]
|
||||||
|
|
||||||
|
wantFullLine(t, full)
|
||||||
|
wantBareLine(t, bare)
|
||||||
|
|
||||||
|
for _, line := range []logLine{full, bare} {
|
||||||
|
got := <-received
|
||||||
|
if got != line.RequestID {
|
||||||
|
t.Errorf("the app received request id %q, the log line has %q",
|
||||||
|
got, line.RequestID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(out.text(), "secret") {
|
||||||
|
t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s",
|
||||||
|
out.text())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantFullLine checks the log line of the request with every header the
|
||||||
|
// line looks at. Its timings are checked by TestTimingsAreInOrder.
|
||||||
|
func wantFullLine(t *testing.T, line logLine) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
headers := map[string]string{"accept": "text/html", "x-custom": "one, two"}
|
||||||
|
|
||||||
|
want := withTimings(line, requestlog.Line{
|
||||||
|
Type: requestType, Time: line.Time, Instance: instance,
|
||||||
|
ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure,
|
||||||
|
Host: appHost, Path: "/full", Protocol: protocol,
|
||||||
|
Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5,
|
||||||
|
RequestID: "from-traefik", PeerIP: localhost,
|
||||||
|
ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group,
|
||||||
|
ContentType: "application/x-www-form-urlencoded", ContentLength: 3,
|
||||||
|
RequestHeaders: headers, HasAuthorization: true, HasCookie: true,
|
||||||
|
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
|
||||||
|
CacheControl: "no-store", Location: "/elsewhere",
|
||||||
|
Action: requestlog.ActionForward,
|
||||||
|
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
||||||
|
})
|
||||||
|
if !reflect.DeepEqual(line.Line, want) {
|
||||||
|
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantBareLine checks the log line of the request with none of them, and
|
||||||
|
// that the fields that do not apply to it are left out.
|
||||||
|
func wantBareLine(t *testing.T, line logLine) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
want := withTimings(line, requestlog.Line{
|
||||||
|
Type: requestType, Time: line.Time, Instance: instance,
|
||||||
|
ClientIP: localhost, Method: http.MethodPost, Scheme: plain,
|
||||||
|
Host: appHost, Path: "/bare", Protocol: protocol,
|
||||||
|
Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID,
|
||||||
|
PeerIP: localhost, ClientGroup: localhost + "/32",
|
||||||
|
UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward,
|
||||||
|
})
|
||||||
|
if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" {
|
||||||
|
t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range []string{
|
||||||
|
"forwarded_for", "content_type", "content_length", "request_headers",
|
||||||
|
"has_authorization", "has_cookie", "websocket", "response_content_type",
|
||||||
|
"cache_control", "location", "counts",
|
||||||
|
} {
|
||||||
|
_, present := line.fields[name]
|
||||||
|
if present {
|
||||||
|
t.Errorf("log line has %s, which does not apply", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// withTimings returns want with the timings of line.
|
||||||
|
func withTimings(line logLine, want requestlog.Line) requestlog.Line {
|
||||||
|
want.DurationTotal = line.DurationTotal
|
||||||
|
want.DurationChecks = line.DurationChecks
|
||||||
|
want.DurationUpstreamConnect = line.DurationUpstreamConnect
|
||||||
|
want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte
|
||||||
|
want.DurationUpstreamTotal = line.DurationUpstreamTotal
|
||||||
|
|
||||||
|
return want
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const hasAuthorization, hasCookie = "has_authorization", "has_cookie"
|
||||||
|
|
||||||
|
for _, tc := range []struct{ header, field, other string }{
|
||||||
|
{"Authorization", hasAuthorization, hasCookie},
|
||||||
|
{"Cookie", hasCookie, hasAuthorization},
|
||||||
|
} {
|
||||||
|
t.Run("only "+tc.header, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||||
|
addr, out := startProxy(t, app.URL, nil)
|
||||||
|
|
||||||
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
|
req.Header.Set(tc.header, "secret")
|
||||||
|
wantStatus(t, do(t, req), http.StatusOK)
|
||||||
|
|
||||||
|
line := out.requestLine(t)
|
||||||
|
|
||||||
|
_, otherPresent := line.fields[tc.other]
|
||||||
|
if line.fields[tc.field] != true || otherPresent {
|
||||||
|
t.Errorf("log line has %s %v and %s %v, want true and none",
|
||||||
|
tc.field, line.fields[tc.field], tc.other, line.fields[tc.other])
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const sentID = "from-traefik"
|
||||||
|
|
||||||
|
sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}}
|
||||||
|
trusted := map[string]string{trustedProxies: trustLocalhost}
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
env map[string]string
|
||||||
|
header http.Header
|
||||||
|
// wantID is the request id logged, "" for a new one.
|
||||||
|
wantID, wantScheme string
|
||||||
|
}{
|
||||||
|
{"a trusted proxy's are kept", trusted, sent, sentID, secure},
|
||||||
|
{"without them, the id is new and the scheme http", trusted, nil, "", plain},
|
||||||
|
{"another peer's are replaced", nil, sent, "", plain},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
received := make(chan string, 2)
|
||||||
|
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
||||||
|
received <- r.Header.Get(requestIDHeader)
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, tc.env)
|
||||||
|
|
||||||
|
// Two requests, so that two new ids can be told apart.
|
||||||
|
ids := make([]string, 0, 2)
|
||||||
|
|
||||||
|
for i := range 2 {
|
||||||
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
|
maps.Copy(req.Header, tc.header)
|
||||||
|
wantStatus(t, do(t, req), http.StatusOK)
|
||||||
|
|
||||||
|
line := out.requestLines(t, i+1)[i]
|
||||||
|
ids = append(ids, line.RequestID)
|
||||||
|
|
||||||
|
got := <-received
|
||||||
|
if line.RequestID != got || line.Scheme != tc.wantScheme {
|
||||||
|
t.Errorf("log line has request_id %q and scheme %q, and the "+
|
||||||
|
"app received id %q; want the same id and scheme %q",
|
||||||
|
line.RequestID, line.Scheme, got, tc.wantScheme)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID):
|
||||||
|
t.Errorf("request ids %q, want %q", ids, tc.wantID)
|
||||||
|
case tc.wantID == "" && (slices.Contains(ids, sentID) ||
|
||||||
|
slices.Contains(ids, "") || ids[0] == ids[1]):
|
||||||
|
t.Errorf("request ids %q, want two new ones", ids)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTimingsAreInOrder(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
||||||
|
|
||||||
|
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
// The pauses set the times apart; a hold-up of the test only
|
||||||
|
// lengthens them.
|
||||||
|
time.Sleep(time.Millisecond)
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_ = http.NewResponseController(w).Flush()
|
||||||
|
|
||||||
|
time.Sleep(time.Millisecond)
|
||||||
|
|
||||||
|
_, _ = io.WriteString(w, "done")
|
||||||
|
})
|
||||||
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
|
trustedProxies: trustLocalhost,
|
||||||
|
denyNets: denied,
|
||||||
|
})
|
||||||
|
|
||||||
|
// Each log line is waited for before the next request, so that the
|
||||||
|
// lines are in order.
|
||||||
|
wantStatus(t, get(t, addr, "/"), http.StatusOK)
|
||||||
|
forwarded := out.requestLines(t, 1)[0]
|
||||||
|
|
||||||
|
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||||
|
req.Header.Set(forwardedFor, denied)
|
||||||
|
wantStatus(t, do(t, req), http.StatusForbidden)
|
||||||
|
refused := out.requestLines(t, 2)[1]
|
||||||
|
|
||||||
|
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
|
||||||
|
health := out.requestLines(t, 3)[2]
|
||||||
|
|
||||||
|
// A request passed to the app has every timing; one refused, none of
|
||||||
|
// the app's; the health check, which runs no check, only the total.
|
||||||
|
wantTimings(t, forwarded, "duration_total", "duration_checks",
|
||||||
|
"duration_upstream_connect", "duration_upstream_first_byte",
|
||||||
|
"duration_upstream_total")
|
||||||
|
wantTimings(t, refused, "duration_total", "duration_checks")
|
||||||
|
wantTimings(t, health, "duration_total")
|
||||||
|
|
||||||
|
if t.Failed() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// In whole microseconds, as they are logged, so that the sum below is
|
||||||
|
// exact.
|
||||||
|
total := microseconds(forwarded.DurationTotal)
|
||||||
|
checks := microseconds(*forwarded.DurationChecks)
|
||||||
|
connect := microseconds(*forwarded.DurationUpstreamConnect)
|
||||||
|
firstByte := microseconds(*forwarded.DurationUpstreamFirstByte)
|
||||||
|
upstream := microseconds(*forwarded.DurationUpstreamTotal)
|
||||||
|
|
||||||
|
// The checks end before the request is handed to the app, and the
|
||||||
|
// connection comes before the answer, which the app ends after a
|
||||||
|
// pause.
|
||||||
|
if checks+upstream > total || connect >= firstByte || firstByte >= upstream {
|
||||||
|
t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+
|
||||||
|
"first byte %d, upstream total %d", total, checks, connect, firstByte,
|
||||||
|
upstream)
|
||||||
|
}
|
||||||
|
|
||||||
|
if *refused.DurationChecks > refused.DurationTotal {
|
||||||
|
t.Errorf("refused request's checks took %v of %v milliseconds",
|
||||||
|
*refused.DurationChecks, refused.DurationTotal)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantTimings checks that the timings named are the only ones line has.
|
||||||
|
func wantTimings(t *testing.T, line logLine, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
var got []string
|
||||||
|
|
||||||
|
for name := range line.fields {
|
||||||
|
if strings.HasPrefix(name, "duration_") {
|
||||||
|
got = append(got, name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
slices.Sort(got)
|
||||||
|
slices.Sort(want)
|
||||||
|
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// microseconds is a timing in whole microseconds.
|
||||||
|
func microseconds(milliseconds float64) int64 {
|
||||||
|
return int64(math.Round(milliseconds * 1000))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLogsAnUpgradedConnection(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, echoAfterUpgrade)
|
||||||
|
addr, out := startProxy(t, app.URL, nil)
|
||||||
|
|
||||||
|
conn := dial(t, addr)
|
||||||
|
send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+
|
||||||
|
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
|
wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols)
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
line := out.requestLine(t)
|
||||||
|
if line.fields["websocket"] != true {
|
||||||
|
t.Errorf("log line has websocket %v, want true", line.fields["websocket"])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,233 @@
|
|||||||
|
package proxy_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testRules are the rules most tests here load: a block rule for
|
||||||
|
// /blocked and a ban rule for /.env.
|
||||||
|
const testRules = `
|
||||||
|
blocked path block ^/blocked$
|
||||||
|
probe path ban ^/\.env$
|
||||||
|
`
|
||||||
|
|
||||||
|
func TestEachRuleAction(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, server := startWithClock(t, "", map[string]string{
|
||||||
|
rulesDir: writeRules(t, "noted path log ^/\n"+testRules),
|
||||||
|
banResponse: "429",
|
||||||
|
})
|
||||||
|
start := clk.Now()
|
||||||
|
|
||||||
|
// A log rule notes its match, and lets the request through.
|
||||||
|
line := s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantRuleIDs(t, line, "noted")
|
||||||
|
|
||||||
|
// A block rule refuses with 403, whatever SWWAF_BAN_RESPONSE is, and
|
||||||
|
// bans no one.
|
||||||
|
line = s.request(client, "/blocked", http.StatusForbidden,
|
||||||
|
requestlog.ActionRuleBlocked)
|
||||||
|
wantRuleIDs(t, line, "noted", "blocked")
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
// A ban rule refuses with SWWAF_BAN_RESPONSE, and bans the client for
|
||||||
|
// seven days, the default.
|
||||||
|
line = s.request(client, "/.env", http.StatusTooManyRequests, requestlog.ActionBanned)
|
||||||
|
wantRuleIDs(t, line, "noted", "probe")
|
||||||
|
|
||||||
|
if line.BanExpires != requestlog.FormatTime(start.Add(7*24*time.Hour)) {
|
||||||
|
t.Errorf("log line has ban_expires %q, want seven days on", line.BanExpires)
|
||||||
|
}
|
||||||
|
|
||||||
|
netblock := netip.MustParsePrefix(client + "/32")
|
||||||
|
want := bans.Ban{
|
||||||
|
Netblock: netblock,
|
||||||
|
Start: start,
|
||||||
|
Expires: start.Add(7 * 24 * time.Hour),
|
||||||
|
Cause: bans.CauseAttack,
|
||||||
|
Reason: "matched the rule probe",
|
||||||
|
Notes: bans.Notes{
|
||||||
|
RuleID: "probe",
|
||||||
|
Target: "path",
|
||||||
|
Request: bans.Request{
|
||||||
|
Time: start,
|
||||||
|
Method: http.MethodGet,
|
||||||
|
Host: appHost,
|
||||||
|
Path: "/.env",
|
||||||
|
Status: http.StatusTooManyRequests,
|
||||||
|
UserAgent: userAgent,
|
||||||
|
},
|
||||||
|
// The four requests up to and including the probe.
|
||||||
|
Requests: 4,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := server.Ledger.Bans(netblock)
|
||||||
|
if len(got) != 1 || got[0] != want {
|
||||||
|
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next request is refused under the ban, without being checked
|
||||||
|
// against the rules, and makes the ban permanent.
|
||||||
|
clk.advance(time.Hour)
|
||||||
|
|
||||||
|
line = s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
|
||||||
|
wantRuleIDs(t, line)
|
||||||
|
|
||||||
|
if line.BanExpires != permanent {
|
||||||
|
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
|
||||||
|
}
|
||||||
|
|
||||||
|
clk.advance(365 * 24 * time.Hour)
|
||||||
|
s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNextClearSignOfAttackAfterABanBansPermanently(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, clk, _ := startWithClock(t, "", map[string]string{
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
attackBanDuration: "1h",
|
||||||
|
})
|
||||||
|
|
||||||
|
// The first probe bans for SWWAF_ATTACK_BAN_DURATION.
|
||||||
|
line := s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
if line.BanExpires != requestlog.FormatTime(clk.Now().Add(time.Hour)) {
|
||||||
|
t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Once that ban has run out without a request, the client is served,
|
||||||
|
// and its next probe bans it for good.
|
||||||
|
clk.advance(time.Hour)
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
|
||||||
|
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
if line.BanExpires != permanent {
|
||||||
|
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRulesComeAfterTheOtherChecks(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
|
||||||
|
exempt = "192.0.2.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||||
|
)
|
||||||
|
|
||||||
|
s, _, server := startWithClock(t, "", map[string]string{
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
allowNets: allowed,
|
||||||
|
rateLimitExemptNets: exempt,
|
||||||
|
rateLimitPerMinute: "1",
|
||||||
|
})
|
||||||
|
|
||||||
|
// A client in SWWAF_ALLOW_NETS is not checked.
|
||||||
|
line := s.request(allowed, "/.env", http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantRuleIDs(t, line)
|
||||||
|
|
||||||
|
// A probe over the rate limit breaks the limit before any rule sees
|
||||||
|
// it.
|
||||||
|
s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionRateLimited)
|
||||||
|
wantRuleIDs(t, line)
|
||||||
|
|
||||||
|
limitBan := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
|
||||||
|
if len(limitBan) != 1 || limitBan[0].Cause != bans.CauseLimit {
|
||||||
|
t.Errorf("bans %+v, want one for a broken limit", limitBan)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A client the rate limits do not apply to is still checked.
|
||||||
|
s.get(exempt, http.StatusOK, requestlog.ActionForward)
|
||||||
|
s.request(exempt, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObserveModeLogsWhatTheRulesWouldDo(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
s, _, server := startWithClock(t, "", map[string]string{
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
mode: observe,
|
||||||
|
})
|
||||||
|
|
||||||
|
line := s.request(client, "/blocked", http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, requestlog.ActionRuleBlocked)
|
||||||
|
wantRuleIDs(t, line, "blocked")
|
||||||
|
|
||||||
|
line = s.request(client, "/.env", http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, requestlog.ActionBanned)
|
||||||
|
wantRuleIDs(t, line, "probe")
|
||||||
|
|
||||||
|
if line.BanExpires != "" {
|
||||||
|
t.Errorf("log line has ban_expires %q, want none", line.BanExpires)
|
||||||
|
}
|
||||||
|
|
||||||
|
// No ban was made.
|
||||||
|
line = s.get(client, http.StatusOK, requestlog.ActionForward)
|
||||||
|
wantWouldAction(t, line, "")
|
||||||
|
|
||||||
|
if got := server.Ledger.Snapshot(); len(got) != 0 {
|
||||||
|
t.Errorf("bans %+v, want none", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const scraper = "192.0.2.200"
|
||||||
|
|
||||||
|
s, _, _ := startWithClock(t, "", map[string]string{
|
||||||
|
rulesDir: writeRules(t, testRules),
|
||||||
|
metricsToken: token,
|
||||||
|
})
|
||||||
|
|
||||||
|
s.request(client, "/blocked", http.StatusForbidden, requestlog.ActionRuleBlocked)
|
||||||
|
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||||
|
|
||||||
|
metrics := s.scrape(scraper)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
|
||||||
|
wantMetric(t, metrics,
|
||||||
|
`smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
|
||||||
|
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
|
||||||
|
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeRules writes content as a rule file into a new directory, and
|
||||||
|
// returns the directory.
|
||||||
|
func writeRules(t *testing.T, content string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
err := os.WriteFile(filepath.Join(dir, "test.rules"), []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write the rule file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return dir
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRuleIDs checks the request log line's rule_ids.
|
||||||
|
func wantRuleIDs(t *testing.T, line logLine, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if !slices.Equal(line.RuleIDs, want) {
|
||||||
|
t.Errorf("log line has rule_ids %v, want %v", line.RuleIDs, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
)
|
||||||
|
|
||||||
|
// checkRules checks the request against the rules of the rule files at
|
||||||
|
// now, notes the ids of those it matches in the log line, and returns the
|
||||||
|
// action of the rule that refuses it, ActionRuleBlocked for a block rule
|
||||||
|
// and ActionBanned for a ban rule, or "" when none does. In enforce mode
|
||||||
|
// a ban rule bans the client's netblock for a clear sign of attack.
|
||||||
|
func (rq *request) checkRules(now time.Time) string {
|
||||||
|
matched := rq.h.rules.Match(rq.in)
|
||||||
|
|
||||||
|
for _, rule := range matched {
|
||||||
|
rq.line.RuleIDs = append(rq.line.RuleIDs, rule.ID)
|
||||||
|
rq.h.metrics.RuleMatched(rule.ID, rule.Action)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(matched) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only the last rule matched can refuse the request.
|
||||||
|
switch last := matched[len(matched)-1]; last.Action {
|
||||||
|
case rules.ActionBlock:
|
||||||
|
return requestlog.ActionRuleBlocked
|
||||||
|
case rules.ActionBan:
|
||||||
|
if !rq.h.config.Observe {
|
||||||
|
rq.banForAttack(now, last)
|
||||||
|
}
|
||||||
|
|
||||||
|
return requestlog.ActionBanned
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -78,7 +78,7 @@ func TestRequestFromAllowNetsIsNotCounted(t *testing.T) {
|
|||||||
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
||||||
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
||||||
{unlistedAddr, 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},
|
||||||
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
|
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
|
||||||
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
|
{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},
|
||||||
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
{listedAddr, http.StatusOK, requestlog.ActionForward},
|
||||||
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
|
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
|
||||||
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
|
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
|
||||||
{fromKP, http.StatusForbidden, requestlog.ActionCountryDenied},
|
{fromKP, http.StatusForbidden, requestlog.ActionCountryDenied},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -28,7 +28,9 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
|
|
||||||
for _, tc := range []struct {
|
for _, tc := range []struct {
|
||||||
name string
|
name string
|
||||||
env map[string]string
|
// limit is the setting set to shortTimeout, which runs out; long
|
||||||
|
// is one set to longTimeoutSetting, which does not, or "".
|
||||||
|
limit, long string
|
||||||
// appTakesNothing has the app never read, while the client sends
|
// appTakesNothing has the app never read, while the client sends
|
||||||
// as fast as it can; otherwise the app reads, and the client
|
// as fast as it can; otherwise the app reads, and the client
|
||||||
// stops sending halfway.
|
// stops sending halfway.
|
||||||
@@ -36,30 +38,26 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
want int
|
want int
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "client request timeout, waiting on the client",
|
name: "client request timeout, waiting on the client",
|
||||||
env: map[string]string{clientRequestTimeout: shortTimeoutSetting},
|
limit: clientRequestTimeout,
|
||||||
want: http.StatusRequestTimeout,
|
want: http.StatusRequestTimeout,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "upstream request timeout, waiting on the client",
|
name: "upstream request timeout, waiting on the client",
|
||||||
env: map[string]string{
|
limit: upstreamRequestTimeout,
|
||||||
upstreamRequestTimeout: shortTimeoutSetting,
|
long: clientRequestTimeout,
|
||||||
clientRequestTimeout: longTimeoutSetting,
|
want: http.StatusRequestTimeout,
|
||||||
},
|
|
||||||
want: http.StatusRequestTimeout,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "upstream request timeout, waiting on the app",
|
name: "upstream request timeout, waiting on the app",
|
||||||
env: map[string]string{upstreamRequestTimeout: shortTimeoutSetting},
|
limit: upstreamRequestTimeout,
|
||||||
appTakesNothing: true,
|
appTakesNothing: true,
|
||||||
want: http.StatusGatewayTimeout,
|
want: http.StatusGatewayTimeout,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "client request timeout, waiting on the app",
|
name: "client request timeout, waiting on the app",
|
||||||
env: map[string]string{
|
limit: clientRequestTimeout,
|
||||||
clientRequestTimeout: shortTimeoutSetting,
|
long: upstreamRequestTimeout,
|
||||||
upstreamRequestTimeout: longTimeoutSetting,
|
|
||||||
},
|
|
||||||
appTakesNothing: true,
|
appTakesNothing: true,
|
||||||
want: http.StatusGatewayTimeout,
|
want: http.StatusGatewayTimeout,
|
||||||
},
|
},
|
||||||
@@ -84,7 +82,12 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
appURL, sendRequest = app.URL, sendPartOfBody
|
appURL, sendRequest = app.URL, sendPartOfBody
|
||||||
}
|
}
|
||||||
|
|
||||||
addr, out := startProxy(t, appURL, tc.env)
|
env := map[string]string{tc.limit: shortTimeoutSetting, metricsToken: token}
|
||||||
|
if tc.long != "" {
|
||||||
|
env[tc.long] = longTimeoutSetting
|
||||||
|
}
|
||||||
|
|
||||||
|
addr, out := startProxy(t, appURL, env)
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
got := readResponse(t, sendRequest(t, addr))
|
got := readResponse(t, sendRequest(t, addr))
|
||||||
wantTimedOut(t, start)
|
wantTimedOut(t, start)
|
||||||
@@ -105,6 +108,7 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
|
|
||||||
wantStatus(t, got, want)
|
wantStatus(t, got, want)
|
||||||
wantLine(t, out.requestLine(t), want, requestlog.ActionTimedOut)
|
wantLine(t, out.requestLine(t), want, requestlog.ActionTimedOut)
|
||||||
|
wantLimitHits(t, addr, tc.limit, 1)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -198,6 +202,7 @@ func TestAppTooSlowToAnswer(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
upstreamResponseTimeout: shortTimeoutSetting,
|
upstreamResponseTimeout: shortTimeoutSetting,
|
||||||
|
metricsToken: token,
|
||||||
})
|
})
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
@@ -213,6 +218,8 @@ func TestAppTooSlowToAnswer(t *testing.T) {
|
|||||||
t.Errorf("log line has upstream_status %v for an app that never answered",
|
t.Errorf("log line has upstream_status %v for an app that never answered",
|
||||||
line.fields["upstream_status"])
|
line.fields["upstream_status"])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
wantLimitHits(t, addr, upstreamResponseTimeout, 1)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
|
func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
|
||||||
@@ -261,6 +268,7 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
clientResponseTimeout: shortTimeoutSetting,
|
clientResponseTimeout: shortTimeoutSetting,
|
||||||
|
metricsToken: token,
|
||||||
})
|
})
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
@@ -272,4 +280,28 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
|
|||||||
line := out.requestLine(t)
|
line := out.requestLine(t)
|
||||||
wantTimedOut(t, start)
|
wantTimedOut(t, start)
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
|
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
|
||||||
|
wantLimitHits(t, addr, clientResponseTimeout, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClosesAnIdleConnection(t *testing.T) {
|
||||||
|
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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,119 @@
|
|||||||
|
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", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
|
||||||
|
{Forwarded: true, Status: 502, ResponseBytes: 12},
|
||||||
|
// Closed without an answer: refused, and no response.
|
||||||
|
{Refused: true, Status: 0},
|
||||||
|
// Answered 404 at smallwebwaf's own endpoints: neither forwarded
|
||||||
|
// nor refused.
|
||||||
|
{Status: 404},
|
||||||
|
} {
|
||||||
|
limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := ratelimit.History{
|
||||||
|
FirstSeen: start,
|
||||||
|
LastSeen: start.Add(6 * time.Minute),
|
||||||
|
Country: "FR",
|
||||||
|
LookedUp: start.Add(3 * time.Minute),
|
||||||
|
Requests: 7,
|
||||||
|
Forwarded: 4,
|
||||||
|
Refused: 2,
|
||||||
|
RequestBytes: 15,
|
||||||
|
ResponseBytes: 122,
|
||||||
|
Responses: ratelimit.Responses{
|
||||||
|
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 2, 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{}
|
||||||
|
}
|
||||||
+304
-47
@@ -1,11 +1,15 @@
|
|||||||
// Package ratelimit counts each client's requests over a minute, an hour
|
// Package ratelimit keeps the table of clients: each client's requests
|
||||||
// and a day, as the "Counting method" section of SPEC.md describes, and
|
// counted over a minute, an hour and a day, as the "Counting method"
|
||||||
// tells when a request takes a client over a rate limit. The counts are
|
// section of SPEC.md describes, which tell when a request takes the client
|
||||||
// kept in memory only, for at most 20,000 clients.
|
// over a rate limit, and each client's history since it was first seen.
|
||||||
|
// At most 20,000 clients are kept, in memory, and written to clients.json
|
||||||
|
// and read from it by the state package.
|
||||||
package ratelimit
|
package ratelimit
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"slices"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -13,7 +17,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// maxClients is how many clients are kept. Past it, the least recently
|
// maxClients is how many clients are kept. Past it, the least recently
|
||||||
// seen client is dropped, and starts afresh if it comes back.
|
// seen client is dropped, with its history, and starts afresh if it comes
|
||||||
|
// back.
|
||||||
const maxClients = 20000
|
const maxClients = 20000
|
||||||
|
|
||||||
const day = 24 * time.Hour
|
const day = 24 * time.Hour
|
||||||
@@ -26,20 +31,99 @@ type Limits struct {
|
|||||||
PerDay int64
|
PerDay int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// Limiter counts each client's requests against the limits. It is safe
|
// Limiter counts each client's requests against the limits, and keeps
|
||||||
// for concurrent use.
|
// its history. It is safe for concurrent use.
|
||||||
type Limiter struct {
|
type Limiter struct {
|
||||||
|
// windows are the minute, the hour and the day, in the order of
|
||||||
|
// Client.buckets.
|
||||||
windows [3]window
|
windows [3]window
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
// clients holds each client's buckets, one pair for each of windows,
|
clients *simplelru.LRU[netip.Prefix, *Client]
|
||||||
// in the same order.
|
}
|
||||||
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
|
|
||||||
|
// Client is a client in the table, as clients.json holds it: its buckets
|
||||||
|
// in each window, and its history.
|
||||||
|
type Client struct {
|
||||||
|
Client netip.Prefix `json:"client"`
|
||||||
|
Minute Buckets `json:"minute"`
|
||||||
|
Hour Buckets `json:"hour"`
|
||||||
|
Day Buckets `json:"day"`
|
||||||
|
History History `json:"history"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Buckets are a client's two buckets in one window: the requests in the
|
||||||
|
// bucket under way, which began at Start, and in the bucket before it.
|
||||||
|
type Buckets struct {
|
||||||
|
Start time.Time `json:"start"`
|
||||||
|
Current int64 `json:"current"`
|
||||||
|
Previous int64 `json:"previous"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// History is what is known of a client since it was first seen.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||||
|
type History struct {
|
||||||
|
FirstSeen time.Time `json:"first_seen"`
|
||||||
|
LastSeen time.Time `json:"last_seen"`
|
||||||
|
// Country is the client's country as it was last looked up, and
|
||||||
|
// LookedUp when that was; both are empty while it never was.
|
||||||
|
Country string `json:"country,omitempty"`
|
||||||
|
LookedUp time.Time `json:"looked_up,omitzero"`
|
||||||
|
// Requests are all the client's requests: Forwarded those passed to
|
||||||
|
// the app, Refused those refused before anything reached it, a 401 at
|
||||||
|
// smallwebwaf's own endpoints included, and neither the others
|
||||||
|
// smallwebwaf answered there.
|
||||||
|
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, Refused for one
|
||||||
|
// refused before anything reached it, a 401 at smallwebwaf's own
|
||||||
|
// endpoints included. Both are false for any other request smallwebwaf
|
||||||
|
// answered there.
|
||||||
|
Forwarded bool
|
||||||
|
Refused bool
|
||||||
|
// Status is what the client was sent, 0 if nothing was.
|
||||||
|
Status int
|
||||||
|
// RequestBytes and ResponseBytes are the body bytes of the request
|
||||||
|
// and of its response.
|
||||||
|
RequestBytes int64
|
||||||
|
ResponseBytes int64
|
||||||
|
// BrokeLimit is true for a request that broke a rate limit.
|
||||||
|
BrokeLimit bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns a Limiter for limits, with no client counted yet.
|
// New returns a Limiter for limits, with no client counted yet.
|
||||||
func New(limits Limits) *Limiter {
|
func New(limits Limits) *Limiter {
|
||||||
clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](maxClients, nil)
|
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
@@ -54,30 +138,194 @@ 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// Counts are a client's requests in the minute, the hour and the day that
|
||||||
|
// end at a request, that request included.
|
||||||
|
type Counts struct {
|
||||||
|
Minute float64 `json:"minute"`
|
||||||
|
Hour float64 `json:"hour"`
|
||||||
|
Day float64 `json:"day"`
|
||||||
|
}
|
||||||
|
|
||||||
// Count counts a request from client at now, in every window, whether or
|
// Count counts a request from client at now, in every window, whether or
|
||||||
// not it is refused. It returns the window whose limit the request takes
|
// not it is refused, and returns the client's requests in each window. It
|
||||||
// the client over, "minute", "hour" or "day", the shortest if it is over
|
// reports whether the request takes the client over a limit, and the
|
||||||
// several, or "" if it is within every limit.
|
// window whose limit it goes over, the shortest if it is over several.
|
||||||
func (l *Limiter) Count(client netip.Prefix, now time.Time) string {
|
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
counts, seen := l.clients.Get(client)
|
var (
|
||||||
if !seen {
|
requests [3]float64
|
||||||
counts = &[3]buckets{}
|
hit Hit
|
||||||
l.clients.Add(client, counts)
|
)
|
||||||
}
|
|
||||||
|
|
||||||
limitHit := ""
|
for i, b := range l.get(client).buckets() {
|
||||||
|
w := l.windows[i]
|
||||||
|
|
||||||
for i, w := range l.windows {
|
requests[i] = b.add(now, w.length)
|
||||||
requests := counts[i].add(now, w.length)
|
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
|
||||||
if limitHit == "" && w.limit > 0 && requests > float64(w.limit) {
|
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
|
||||||
limitHit = w.name
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return limitHit
|
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
|
||||||
|
|
||||||
|
return counts, hit, hit.Window != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset sets client's counts in every window back to zero. Its history
|
||||||
|
// 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++
|
||||||
|
}
|
||||||
|
|
||||||
|
if r.Refused {
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// Len returns how many clients are in the table.
|
||||||
|
func (l *Limiter) Len() int {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
return l.clients.Len()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Snapshot returns every client in the table, sorted by address, as
|
||||||
|
// 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 the table, in place of
|
||||||
|
// the clients it holds, 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) {
|
||||||
|
clients = slices.Clone(clients)
|
||||||
|
slices.SortStableFunc(clients, func(a, b Client) int {
|
||||||
|
return a.History.LastSeen.Compare(b.History.LastSeen)
|
||||||
|
})
|
||||||
|
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
l.clients.Purge()
|
||||||
|
|
||||||
|
for _, c := range clients {
|
||||||
|
for i, b := range c.buckets() {
|
||||||
|
// The window that ends at now covers neither bucket once it
|
||||||
|
// begins after the bucket under way has ended.
|
||||||
|
length := l.windows[i].length
|
||||||
|
if !now.Add(-length).Before(b.Start.Add(length)) {
|
||||||
|
*b = Buckets{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
l.clients.Add(c.Client, &c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// get returns client's entry in the table, a new one if it has none, and
|
||||||
|
// makes it the most recently seen.
|
||||||
|
func (l *Limiter) get(client netip.Prefix) *Client {
|
||||||
|
c, seen := l.clients.Get(client)
|
||||||
|
if !seen {
|
||||||
|
c = &Client{Client: client}
|
||||||
|
l.clients.Add(client, c)
|
||||||
|
}
|
||||||
|
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// buckets returns c's buckets in the minute, the hour and the day.
|
||||||
|
func (c *Client) buckets() [3]*Buckets {
|
||||||
|
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
|
||||||
}
|
}
|
||||||
|
|
||||||
// window is a length of time over which requests are counted, and the
|
// window is a length of time over which requests are counted, and the
|
||||||
@@ -88,14 +336,6 @@ type window struct {
|
|||||||
limit int64
|
limit int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// buckets are a client's two buckets in one window: the requests in the
|
|
||||||
// bucket under way, which began at start, and in the bucket before it.
|
|
||||||
type buckets struct {
|
|
||||||
start time.Time
|
|
||||||
current int64
|
|
||||||
previous int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// add counts a request at now in a window of length, and returns the
|
// add counts a request at now in a window of length, and returns the
|
||||||
// client's requests in the window that ends at now: those in the bucket
|
// client's requests in the window that ends at now: those in the bucket
|
||||||
// under way, and those in the bucket before it weighted by how much of
|
// under way, and those in the bucket before it weighted by how much of
|
||||||
@@ -106,27 +346,44 @@ type buckets struct {
|
|||||||
// bucket. A request dated more than a second before it means the clock
|
// bucket. A request dated more than a second before it means the clock
|
||||||
// was set back, and the buckets start afresh: otherwise the bucket before
|
// was set back, and the buckets start afresh: otherwise the bucket before
|
||||||
// would keep its full weight until the clock caught up.
|
// would keep its full weight until the clock caught up.
|
||||||
func (b *buckets) add(now time.Time, length time.Duration) float64 {
|
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
||||||
if now.Before(b.start.Add(-time.Second)) {
|
if now.Before(b.Start.Add(-time.Second)) {
|
||||||
*b = buckets{}
|
*b = Buckets{}
|
||||||
}
|
}
|
||||||
|
|
||||||
start := now.Truncate(length)
|
start := now.Truncate(length)
|
||||||
if start.After(b.start) {
|
if start.After(b.Start) {
|
||||||
if start.Equal(b.start.Add(length)) {
|
if start.Equal(b.Start.Add(length)) {
|
||||||
b.previous = b.current
|
b.Previous = b.Current
|
||||||
} else {
|
} else {
|
||||||
b.previous = 0
|
b.Previous = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.start = start
|
b.Start = start
|
||||||
b.current = 0
|
b.Current = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.current++
|
b.Current++
|
||||||
|
|
||||||
elapsed := max(now.Sub(b.start), 0)
|
elapsed := max(now.Sub(b.Start), 0)
|
||||||
covered := 1 - float64(elapsed)/float64(length)
|
covered := 1 - float64(elapsed)/float64(length)
|
||||||
|
|
||||||
return float64(b.previous)*covered + float64(b.current)
|
return float64(b.Previous)*covered + float64(b.Current)
|
||||||
|
}
|
||||||
|
|
||||||
|
// add counts a response with status in its class. A status of 0, for
|
||||||
|
// nothing sent, is not a response.
|
||||||
|
func (r *Responses) add(status int) {
|
||||||
|
switch {
|
||||||
|
case status >= http.StatusInternalServerError:
|
||||||
|
r.Status5xx++
|
||||||
|
case status >= http.StatusBadRequest:
|
||||||
|
r.Status4xx++
|
||||||
|
case status >= http.StatusMultipleChoices:
|
||||||
|
r.Status3xx++
|
||||||
|
case status >= http.StatusOK:
|
||||||
|
r.Status2xx++
|
||||||
|
case status >= http.StatusContinue:
|
||||||
|
r.Status1xx++
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,6 +54,75 @@ 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 TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
limiter := ratelimit.New(ratelimit.Limits{})
|
||||||
|
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||||
|
start := midnight()
|
||||||
|
|
||||||
|
for range 3 {
|
||||||
|
limiter.Count(client, start)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A quarter into the next hour, the minute has only this request. The
|
||||||
|
// hour still covers three quarters of the bucket before, with its three
|
||||||
|
// requests, which count 2.25, and this one: 3.25. The day covers all
|
||||||
|
// four.
|
||||||
|
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
|
||||||
|
|
||||||
|
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
|
||||||
|
if counts != want {
|
||||||
|
t.Errorf("counts %+v, want %+v", counts, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResetSetsTheCountsBackToZero(t *testing.T) {
|
||||||
|
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) {
|
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -192,9 +261,9 @@ func wantCount(
|
|||||||
) {
|
) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
got := limiter.Count(client, now)
|
_, hit, _ := limiter.Count(client, now)
|
||||||
if got != want {
|
if hit.Window != want {
|
||||||
t.Errorf("request from %s at %s is over %q, want %q",
|
t.Errorf("request from %s at %s is over %q, want %q",
|
||||||
client, now.Format(time.RFC3339), got, want)
|
client, now.Format(time.RFC3339), hit.Window, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,310 @@
|
|||||||
|
// Package remotelog sends the lines smallwebwaf writes on stdout to the
|
||||||
|
// remote log endpoint, SWWAF_LOG_REMOTE_URL, as the "Request log" section
|
||||||
|
// of SPEC.md describes: each line as the message of an RFC 5424 syslog
|
||||||
|
// record, over UDP, TCP or TLS. Lines wait in a bounded buffer, so a slow
|
||||||
|
// or unreachable endpoint never holds up a request or stdout.
|
||||||
|
package remotelog
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
|
"crypto/x509"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"sync/atomic"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The forms of SWWAF_LOG_REMOTE_URL, by its scheme.
|
||||||
|
const (
|
||||||
|
SchemeUDP = "syslog+udp"
|
||||||
|
SchemeTCP = "syslog+tcp"
|
||||||
|
SchemeTLS = "syslog+tls"
|
||||||
|
)
|
||||||
|
|
||||||
|
// A record's priority is the number of its facility times the number of
|
||||||
|
// severities there are, plus the number of its severity. Every record's
|
||||||
|
// severity is informational.
|
||||||
|
const (
|
||||||
|
severities = 8
|
||||||
|
informational = 6
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// dialTimeout bounds connecting to the endpoint, the TLS handshake
|
||||||
|
// included.
|
||||||
|
dialTimeout = 10 * time.Second
|
||||||
|
// After a failed attempt to connect, or a connection on which a record
|
||||||
|
// fails, the next attempt to connect is made a second later, and
|
||||||
|
// retryDelayFactor times as long after each further failure in a row,
|
||||||
|
// up to a minute. A connection that fails after it has stayed up for
|
||||||
|
// resetRetryDelayAfter ends the row.
|
||||||
|
firstRetryDelay = time.Second
|
||||||
|
retryDelayFactor = 2
|
||||||
|
maxRetryDelay = time.Minute
|
||||||
|
resetRetryDelayAfter = time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
// Params are what New needs.
|
||||||
|
type Params struct {
|
||||||
|
// URL is the endpoint (SWWAF_LOG_REMOTE_URL): SchemeUDP, SchemeTCP or
|
||||||
|
// SchemeTLS, a host and a port.
|
||||||
|
URL *url.URL
|
||||||
|
// RootCAs are the certificates a SchemeTLS endpoint's certificate
|
||||||
|
// must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil for the host's.
|
||||||
|
RootCAs *x509.CertPool
|
||||||
|
// Buffer is the most lines held while they wait to be sent
|
||||||
|
// (SWWAF_LOG_REMOTE_BUFFER).
|
||||||
|
Buffer int
|
||||||
|
// Facility is the number of the records' syslog facility
|
||||||
|
// (SWWAF_LOG_REMOTE_FACILITY), and AppName their APP-NAME
|
||||||
|
// (SWWAF_LOG_REMOTE_APP_NAME).
|
||||||
|
Facility int
|
||||||
|
AppName string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sender sends lines to the endpoint. Write puts them in its buffer, and
|
||||||
|
// Run sends them from there.
|
||||||
|
type Sender struct {
|
||||||
|
url *url.URL
|
||||||
|
tlsConfig *tls.Config
|
||||||
|
// beforeTime and afterTime are the parts of every record's header
|
||||||
|
// before and after its time, as RFC 5424 lays the header out.
|
||||||
|
beforeTime string
|
||||||
|
afterTime string
|
||||||
|
// records is the buffer: each line's record, framed to be sent.
|
||||||
|
records chan []byte
|
||||||
|
sent atomic.Int64
|
||||||
|
dropped atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// New returns a Sender for the endpoint params.URL.
|
||||||
|
func New(params Params) *Sender {
|
||||||
|
hostname, err := os.Hostname()
|
||||||
|
if err != nil || hostname == "" {
|
||||||
|
hostname = "-" // RFC 5424's value for a field that has none
|
||||||
|
}
|
||||||
|
|
||||||
|
priority := params.Facility*severities + informational
|
||||||
|
|
||||||
|
return &Sender{
|
||||||
|
url: params.URL,
|
||||||
|
tlsConfig: &tls.Config{
|
||||||
|
RootCAs: params.RootCAs,
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
},
|
||||||
|
// The 1 is the version of the format. The process id, the message
|
||||||
|
// id and the structured data have no value.
|
||||||
|
beforeTime: "<" + strconv.Itoa(priority) + ">1 ",
|
||||||
|
afterTime: " " + hostname + " " + params.AppName + " - - - ",
|
||||||
|
records: make(chan []byte, params.Buffer),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write puts each line in p in the buffer, as the message of a record of
|
||||||
|
// its own, and never waits: when the buffer is full, the oldest record in
|
||||||
|
// it is dropped to make room. It is safe for concurrent use.
|
||||||
|
func (s *Sender) Write(p []byte) (int, error) {
|
||||||
|
at := requestlog.FormatTime(time.Now())
|
||||||
|
|
||||||
|
for line := range bytes.Lines(p) {
|
||||||
|
line = bytes.TrimSuffix(line, []byte("\n"))
|
||||||
|
if len(line) > 0 {
|
||||||
|
s.put(s.record(at, line))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sent is how many records have been sent.
|
||||||
|
func (s *Sender) Sent() int64 {
|
||||||
|
return s.sent.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Dropped is how many records were dropped: the oldest in a full buffer,
|
||||||
|
// and those whose sending failed.
|
||||||
|
func (s *Sender) Dropped() int64 {
|
||||||
|
return s.dropped.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Depth is how many records are in the buffer.
|
||||||
|
func (s *Sender) Depth() int {
|
||||||
|
return len(s.records)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run connects to the endpoint and sends each record as it comes into the
|
||||||
|
// buffer, until ctx is done. Then it sends the records still in the buffer,
|
||||||
|
// on the connection open at that time or, if there is none, on a new one,
|
||||||
|
// until none is left or one fails, and returns. How long it may take over
|
||||||
|
// that is for the caller to bound.
|
||||||
|
//
|
||||||
|
// A connection on which a record fails is closed and the record dropped.
|
||||||
|
// That failure, like a failed attempt to connect, is logged to processLog
|
||||||
|
// and followed by the next attempt after firstRetryDelay, retryDelayFactor
|
||||||
|
// times as long after each further failure in a row up to maxRetryDelay,
|
||||||
|
// and firstRetryDelay again after a connection that stayed up for
|
||||||
|
// resetRetryDelayAfter. Meanwhile the records wait in the buffer.
|
||||||
|
func (s *Sender) Run(ctx context.Context, processLog *slog.Logger) {
|
||||||
|
conn := s.send(ctx, processLog)
|
||||||
|
if conn == nil && len(s.records) > 0 {
|
||||||
|
conn, _ = s.dial(context.WithoutCancel(ctx))
|
||||||
|
}
|
||||||
|
|
||||||
|
if conn == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = conn.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case record := <-s.records:
|
||||||
|
if s.write(conn, record) != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// record returns line as an RFC 5424 record made at the time at, framed
|
||||||
|
// for the endpoint: on its own over UDP, since each datagram holds one,
|
||||||
|
// and over TCP and TLS after its length in bytes and a space, the
|
||||||
|
// octet-counted framing of RFC 6587 and RFC 5425.
|
||||||
|
func (s *Sender) record(at string, line []byte) []byte {
|
||||||
|
record := make([]byte, 0, len(s.beforeTime)+len(at)+len(s.afterTime)+len(line))
|
||||||
|
record = append(record, s.beforeTime...)
|
||||||
|
record = append(record, at...)
|
||||||
|
record = append(record, s.afterTime...)
|
||||||
|
record = append(record, line...)
|
||||||
|
|
||||||
|
if s.url.Scheme == SchemeUDP {
|
||||||
|
return record
|
||||||
|
}
|
||||||
|
|
||||||
|
return append([]byte(strconv.Itoa(len(record))+" "), record...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// put adds record to the buffer, first dropping the oldest record in it
|
||||||
|
// while it is full.
|
||||||
|
func (s *Sender) put(record []byte) {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case s.records <- record:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-s.records:
|
||||||
|
s.dropped.Add(1)
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// send connects to the endpoint and sends each record as it comes into
|
||||||
|
// the buffer, until ctx is done, and returns the connection then open, or
|
||||||
|
// nil.
|
||||||
|
func (s *Sender) send(ctx context.Context, processLog *slog.Logger) net.Conn {
|
||||||
|
delay := firstRetryDelay
|
||||||
|
|
||||||
|
for {
|
||||||
|
conn, err := s.dial(ctx)
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
connected := time.Now()
|
||||||
|
|
||||||
|
err = s.sendOn(ctx, conn)
|
||||||
|
if err == nil {
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
if time.Since(connected) >= resetRetryDelayAfter {
|
||||||
|
delay = firstRetryDelay
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
processLog.Warn("sending to SWWAF_LOG_REMOTE_URL failed",
|
||||||
|
"error", err.Error(), "connecting_again_in", delay.String())
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-time.After(delay):
|
||||||
|
case <-ctx.Done():
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
delay = min(retryDelayFactor*delay, maxRetryDelay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendOn sends each record on conn as it comes into the buffer, until one
|
||||||
|
// fails, whose error it returns, or ctx is done.
|
||||||
|
func (s *Sender) sendOn(ctx context.Context, conn net.Conn) error {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case record := <-s.records:
|
||||||
|
err := s.write(conn, record)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
case <-ctx.Done():
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// write sends record on conn, and counts it as sent or, if that fails,
|
||||||
|
// as dropped. A record too long for one UDP datagram is dropped without
|
||||||
|
// an error, since the connection has not failed: a long request must not
|
||||||
|
// hold up the lines after it.
|
||||||
|
func (s *Sender) write(conn net.Conn, record []byte) error {
|
||||||
|
_, err := conn.Write(record)
|
||||||
|
if err != nil {
|
||||||
|
s.dropped.Add(1)
|
||||||
|
|
||||||
|
if errors.Is(err, syscall.EMSGSIZE) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Errorf("send a record: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.sent.Add(1)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dial connects to the endpoint.
|
||||||
|
func (s *Sender) dial(ctx context.Context) (net.Conn, error) {
|
||||||
|
dialer := &net.Dialer{Timeout: dialTimeout}
|
||||||
|
|
||||||
|
switch s.url.Scheme {
|
||||||
|
case SchemeUDP:
|
||||||
|
return dialer.DialContext(ctx, "udp", s.url.Host)
|
||||||
|
case SchemeTLS:
|
||||||
|
tlsDialer := &tls.Dialer{NetDialer: dialer, Config: s.tlsConfig}
|
||||||
|
|
||||||
|
return tlsDialer.DialContext(ctx, "tcp", s.url.Host)
|
||||||
|
default:
|
||||||
|
return dialer.DialContext(ctx, "tcp", s.url.Host)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,640 @@
|
|||||||
|
package remotelog_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/ecdsa"
|
||||||
|
"crypto/elliptic"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/tls"
|
||||||
|
"crypto/x509"
|
||||||
|
"crypto/x509/pkix"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"math/big"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The tests run in a synctest bubble, where the time package runs on a
|
||||||
|
// clock of the test's own, which starts at 2000-01-01T00:00:00Z: a wait
|
||||||
|
// lasts exactly as long as it should, however slowly the test process
|
||||||
|
// runs, and synctest.Wait returns once the sender has done all it can
|
||||||
|
// before time passes. The endpoint is a listener on the loopback address.
|
||||||
|
// A test reads from it only once the records are on their way, and checks
|
||||||
|
// the sender's counts first, since a goroutine of the bubble that waits on
|
||||||
|
// the network keeps that clock from moving on. For the same reason the
|
||||||
|
// endpoint that refuses connections, a tlsEndpoint, runs outside the
|
||||||
|
// bubble: a sender connecting over TLS waits on the endpoint's answer.
|
||||||
|
|
||||||
|
const (
|
||||||
|
// started is the time a record made as a test starts gives.
|
||||||
|
started = "2000-01-01T00:00:00.000Z"
|
||||||
|
appName = "fsn1app1/gitea"
|
||||||
|
// local0 is the number of the default facility, and local0Info the
|
||||||
|
// priority of its records.
|
||||||
|
local0 = 16
|
||||||
|
local0Info = "<134>"
|
||||||
|
// loopback is where the endpoints listen.
|
||||||
|
loopback = "127.0.0.1:0"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRecordsOverUDPGoOnePerDatagram(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = endpoint.Close() })
|
||||||
|
|
||||||
|
sender, _, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte(`{"type":"request"}` + "\n" + `{"type":"process"}` + "\n"))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, 2, 0, 0)
|
||||||
|
|
||||||
|
for _, line := range []string{`{"type":"request"}`, `{"type":"process"}`} {
|
||||||
|
datagram := make([]byte, 1024)
|
||||||
|
|
||||||
|
n, _, err := endpoint.ReadFrom(datagram)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := record(t, local0Info, appName, line)
|
||||||
|
if string(datagram[:n]) != want {
|
||||||
|
t.Errorf("datagram %q, want %q", datagram[:n], want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordsOverTCPAreOctetCountedWithTheirFacility(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpoint := listen(t)
|
||||||
|
endpointParams := params(remotelog.SchemeTCP, endpoint.Addr())
|
||||||
|
endpointParams.Facility = 19 // local3
|
||||||
|
endpointParams.AppName = "gitea"
|
||||||
|
sender, _, _ := run(t, endpointParams)
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte("first\nsecond\n"))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, 2, 0, 0)
|
||||||
|
|
||||||
|
frames := bufio.NewReader(accept(t, endpoint))
|
||||||
|
wantFrame(t, frames, record(t, "<158>", "gitea", "first"))
|
||||||
|
wantFrame(t, frames, record(t, "<158>", "gitea", "second"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStalledEndpointHoldsUpNoWriteAndOldestRecordsAreDropped(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
certificate, roots := testCertificate(t)
|
||||||
|
endpoint := listen(t)
|
||||||
|
endpointParams := params(remotelog.SchemeTLS, endpoint.Addr())
|
||||||
|
endpointParams.RootCAs = roots
|
||||||
|
endpointParams.Buffer = 3
|
||||||
|
sender, _, _ := run(t, endpointParams)
|
||||||
|
|
||||||
|
// The sender connects, and its TLS handshake waits for an answer
|
||||||
|
// the endpoint does not give yet.
|
||||||
|
conn := accept(t, endpoint)
|
||||||
|
|
||||||
|
var stdout bytes.Buffer
|
||||||
|
|
||||||
|
out := io.MultiWriter(&stdout, sender)
|
||||||
|
for i := range 5 {
|
||||||
|
_, _ = fmt.Fprintf(out, "line %d\n", i+1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if stdout.String() != "line 1\nline 2\nline 3\nline 4\nline 5\n" {
|
||||||
|
t.Errorf("stdout has %q", stdout.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
wantCounts(t, sender, 0, 2, 3)
|
||||||
|
|
||||||
|
// Once the endpoint answers, the three newest records are sent.
|
||||||
|
server := tls.Server(conn, &tls.Config{
|
||||||
|
Certificates: []tls.Certificate{certificate},
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
})
|
||||||
|
|
||||||
|
err := server.HandshakeContext(t.Context())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("handshake: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, 3, 2, 0)
|
||||||
|
|
||||||
|
frames := bufio.NewReader(server)
|
||||||
|
for _, line := range []string{"line 3", "line 4", "line 5"} {
|
||||||
|
wantFrame(t, frames, record(t, local0Info, appName, line))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReconnectsWithBackoffAfterTheEndpointGoesAway(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
certificate, roots := testCertificate(t)
|
||||||
|
endpoint := startTLSEndpoint(t, certificate)
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
|
||||||
|
endpointParams.RootCAs = roots
|
||||||
|
sender, logged, _ := run(t, endpointParams)
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte("one\n"))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, 1, 0, 0)
|
||||||
|
|
||||||
|
conn := endpoint.next(t)
|
||||||
|
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "one"))
|
||||||
|
|
||||||
|
// The endpoint goes away: it closes the connection, and refuses the
|
||||||
|
// next ones. The sender notices when a record fails, and tries to
|
||||||
|
// connect again a second later, then two seconds after that.
|
||||||
|
endpoint.refusing.Store(true)
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
writeUntilDropped(t, sender, 1)
|
||||||
|
sent := sender.Sent()
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte("two\n"))
|
||||||
|
|
||||||
|
time.Sleep(time.Second)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
endpoint.refusing.Store(false)
|
||||||
|
|
||||||
|
time.Sleep(2*time.Second - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, sent, 1, 1)
|
||||||
|
|
||||||
|
// The endpoint is back, and the record waiting is sent.
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, sent+1, 1, 0)
|
||||||
|
|
||||||
|
conn = endpoint.next(t)
|
||||||
|
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "two"))
|
||||||
|
wantRetries(t, logged, "1s", "2s")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAConnectionClosedAtOnceIsMadeAgainAfterAGrowingDelay(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpoint := listen(t)
|
||||||
|
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
|
||||||
|
|
||||||
|
// The endpoint closes each connection as soon as it takes it. The
|
||||||
|
// sender notices when a record fails, and connects again a second
|
||||||
|
// later, then two seconds after that, then four.
|
||||||
|
delays := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second}
|
||||||
|
for i, delay := range delays {
|
||||||
|
_ = accept(t, endpoint).Close()
|
||||||
|
|
||||||
|
writeUntilDropped(t, sender, int64(i+1))
|
||||||
|
wantConnectedAgainAfter(t, sender, delay)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantRetries(t, logged, "1s", "2s", "4s")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTheDelayStartsAgainAfterAConnectionThatStayedUpAMinute(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpoint := listen(t)
|
||||||
|
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
|
||||||
|
|
||||||
|
_ = accept(t, endpoint).Close()
|
||||||
|
|
||||||
|
writeUntilDropped(t, sender, 1)
|
||||||
|
wantConnectedAgainAfter(t, sender, time.Second)
|
||||||
|
|
||||||
|
// A connection that fails just short of a minute after it was made
|
||||||
|
// leaves the delay growing.
|
||||||
|
conn := accept(t, endpoint)
|
||||||
|
|
||||||
|
time.Sleep(time.Minute - time.Nanosecond)
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
writeUntilDropped(t, sender, 2)
|
||||||
|
wantConnectedAgainAfter(t, sender, 2*time.Second)
|
||||||
|
|
||||||
|
// One that fails a minute after it was made starts it again from a
|
||||||
|
// second.
|
||||||
|
conn = accept(t, endpoint)
|
||||||
|
|
||||||
|
time.Sleep(time.Minute)
|
||||||
|
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
writeUntilDropped(t, sender, 3)
|
||||||
|
wantConnectedAgainAfter(t, sender, time.Second)
|
||||||
|
wantRetries(t, logged, "1s", "2s", "1s")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestALineTooLongForADatagramIsDroppedAlone(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = endpoint.Close() })
|
||||||
|
|
||||||
|
sender, logged, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
|
||||||
|
|
||||||
|
// With its header, the first line's record is longer than the 65507
|
||||||
|
// bytes a UDP datagram over IPv4 holds. It is dropped, nothing is
|
||||||
|
// logged, and the next line is sent at once.
|
||||||
|
_, _ = sender.Write([]byte(strings.Repeat("x", 65507) + "\nnext\n"))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantCounts(t, sender, 1, 1, 0)
|
||||||
|
wantRetries(t, logged)
|
||||||
|
|
||||||
|
datagram := make([]byte, 1024)
|
||||||
|
|
||||||
|
n, _, err := endpoint.ReadFrom(datagram)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := record(t, local0Info, appName, "next")
|
||||||
|
if string(datagram[:n]) != want {
|
||||||
|
t.Errorf("datagram %q, want %q", datagram[:n], want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordsWaitingAtTheStopAreSent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
certificate, roots := testCertificate(t)
|
||||||
|
endpoint := startTLSEndpoint(t, certificate)
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
// The endpoint refuses the sender's first connection: it fails to
|
||||||
|
// connect, and waits a second to try again.
|
||||||
|
endpoint.refusing.Store(true)
|
||||||
|
|
||||||
|
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
|
||||||
|
endpointParams.RootCAs = roots
|
||||||
|
sender, logged, stop := run(t, endpointParams)
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
wantRetries(t, logged, "1s")
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte("one\ntwo\n"))
|
||||||
|
|
||||||
|
endpoint.refusing.Store(false)
|
||||||
|
|
||||||
|
// Stopped before that second is over, it connects to send them.
|
||||||
|
stop()
|
||||||
|
wantCounts(t, sender, 2, 0, 0)
|
||||||
|
|
||||||
|
frames := bufio.NewReader(endpoint.next(t))
|
||||||
|
wantFrame(t, frames, record(t, local0Info, appName, "one"))
|
||||||
|
wantFrame(t, frames, record(t, local0Info, appName, "two"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// output collects what the sender logs.
|
||||||
|
type output struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
buf bytes.Buffer
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write adds lines the sender logs.
|
||||||
|
func (o *output) Write(p []byte) (int, error) {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
return o.buf.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
// text returns everything logged so far.
|
||||||
|
func (o *output) text() string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
return o.buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// params returns the settings of a Sender for the endpoint at addr, in
|
||||||
|
// the form scheme names: room for ten lines, the default facility, and
|
||||||
|
// appName.
|
||||||
|
func params(scheme string, addr net.Addr) remotelog.Params {
|
||||||
|
return remotelog.Params{
|
||||||
|
URL: &url.URL{Scheme: scheme, Host: addr.String()},
|
||||||
|
Buffer: 10,
|
||||||
|
Facility: local0,
|
||||||
|
AppName: appName,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// run runs a Sender with settings until the test ends or the function
|
||||||
|
// it returns is called, which waits for Run to return. It returns the
|
||||||
|
// Sender, and what it logs.
|
||||||
|
func run(t *testing.T, settings remotelog.Params) (*remotelog.Sender, *output, func()) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
sender := remotelog.New(settings)
|
||||||
|
logged := &output{}
|
||||||
|
ctx, cancel := context.WithCancel(t.Context())
|
||||||
|
ran := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
sender.Run(ctx, slog.New(slog.NewJSONHandler(logged, nil)))
|
||||||
|
close(ran)
|
||||||
|
}()
|
||||||
|
|
||||||
|
stop := func() {
|
||||||
|
cancel()
|
||||||
|
<-ran
|
||||||
|
}
|
||||||
|
t.Cleanup(stop)
|
||||||
|
|
||||||
|
return sender, logged, stop
|
||||||
|
}
|
||||||
|
|
||||||
|
// listen returns a TCP listener on the loopback address, closed when the
|
||||||
|
// test ends.
|
||||||
|
func listen(t *testing.T) net.Listener {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", loopback)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = listener.Close() })
|
||||||
|
|
||||||
|
return listener
|
||||||
|
}
|
||||||
|
|
||||||
|
// accept returns the next connection to listener, closed when the test
|
||||||
|
// ends.
|
||||||
|
func accept(t *testing.T, listener net.Listener) net.Conn {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("accept: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = conn.Close() })
|
||||||
|
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// tlsEndpoint is a syslog+tls endpoint on the loopback address, which a
|
||||||
|
// test starts outside its bubble. It keeps its listener until the test
|
||||||
|
// ends, and either takes each connection or refuses it.
|
||||||
|
type tlsEndpoint struct {
|
||||||
|
addr net.Addr
|
||||||
|
// refusing is set while the endpoint closes each connection before the
|
||||||
|
// TLS handshake, which fails the sender's attempt to connect.
|
||||||
|
refusing atomic.Bool
|
||||||
|
// conns are the connections it has taken, after the handshake.
|
||||||
|
conns chan net.Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// startTLSEndpoint starts a tlsEndpoint with certificate, which takes
|
||||||
|
// connections until it is told to refuse them.
|
||||||
|
func startTLSEndpoint(t *testing.T, certificate tls.Certificate) *tlsEndpoint {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
listener := listen(t)
|
||||||
|
endpoint := &tlsEndpoint{addr: listener.Addr(), conns: make(chan net.Conn, 10)}
|
||||||
|
config := &tls.Config{
|
||||||
|
Certificates: []tls.Certificate{certificate},
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
for {
|
||||||
|
conn, err := listener.Accept()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
server := tls.Server(conn, config)
|
||||||
|
if endpoint.refusing.Load() || server.HandshakeContext(t.Context()) != nil {
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
endpoint.conns <- server
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return endpoint
|
||||||
|
}
|
||||||
|
|
||||||
|
// next returns the next connection the endpoint has taken, closed when
|
||||||
|
// the test ends.
|
||||||
|
func (e *tlsEndpoint) next(t *testing.T) net.Conn {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
conn := <-e.conns
|
||||||
|
|
||||||
|
t.Cleanup(func() { _ = conn.Close() })
|
||||||
|
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// record returns the record of line made as the test started, with the
|
||||||
|
// priority and the app name given.
|
||||||
|
func record(t *testing.T, priority, app, line string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
hostname, err := os.Hostname()
|
||||||
|
if err != nil || hostname == "" {
|
||||||
|
hostname = "-"
|
||||||
|
}
|
||||||
|
|
||||||
|
return priority + "1 " + started + " " + hostname + " " + app + " - - - " + line
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantFrame reads the next octet-counted frame from frames, and checks
|
||||||
|
// that it holds want.
|
||||||
|
func wantFrame(t *testing.T, frames *bufio.Reader, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
count, err := frames.ReadString(' ')
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read a frame's length: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
length, err := strconv.Atoi(strings.TrimSuffix(count, " "))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("frame starts %q, not with its length", count)
|
||||||
|
}
|
||||||
|
|
||||||
|
got := make([]byte, length)
|
||||||
|
|
||||||
|
_, err = io.ReadFull(frames, got)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read a frame: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if string(got) != want {
|
||||||
|
t.Errorf("frame %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantCounts checks the records sender has sent, dropped and holds in
|
||||||
|
// its buffer.
|
||||||
|
func wantCounts(
|
||||||
|
t *testing.T, sender *remotelog.Sender, sent, dropped int64, depth int,
|
||||||
|
) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if sender.Sent() != sent || sender.Dropped() != dropped || sender.Depth() != depth {
|
||||||
|
t.Fatalf("sent %d, dropped %d, %d in the buffer; want %d, %d and %d",
|
||||||
|
sender.Sent(), sender.Dropped(), sender.Depth(), sent, dropped, depth)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeUntilDropped writes a line at a time until the count of records
|
||||||
|
// sender has dropped reaches dropped. The records it sends on a
|
||||||
|
// connection the endpoint has closed are lost before one fails; how many
|
||||||
|
// depends on when the endpoint's host answers that the connection is
|
||||||
|
// gone.
|
||||||
|
func writeUntilDropped(t *testing.T, sender *remotelog.Sender, dropped int64) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for sender.Dropped() < dropped {
|
||||||
|
_, _ = sender.Write([]byte("lost\n"))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantConnectedAgainAfter writes a line while the sender waits to connect
|
||||||
|
// again, and checks that it connects, and takes the line from the buffer,
|
||||||
|
// only once delay is over.
|
||||||
|
func wantConnectedAgainAfter(
|
||||||
|
t *testing.T, sender *remotelog.Sender, delay time.Duration,
|
||||||
|
) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
_, _ = sender.Write([]byte("waiting\n"))
|
||||||
|
|
||||||
|
time.Sleep(delay - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
if sender.Depth() != 1 {
|
||||||
|
t.Fatalf("connected again before %v", delay)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
if sender.Depth() != 0 {
|
||||||
|
t.Fatalf("not connected again after %v", delay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRetries checks that the sender logged a failure, of an attempt to
|
||||||
|
// connect or of a connection, for each of delays, the time until the next
|
||||||
|
// attempt, in order, and logged nothing else.
|
||||||
|
func wantRetries(t *testing.T, logged *output, delays ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
var got []string
|
||||||
|
|
||||||
|
for line := range strings.Lines(logged.text()) {
|
||||||
|
var fields map[string]any
|
||||||
|
|
||||||
|
err := json.Unmarshal([]byte(line), &fields)
|
||||||
|
if err != nil || fields["msg"] != "sending to SWWAF_LOG_REMOTE_URL failed" {
|
||||||
|
t.Fatalf("logged %q", line)
|
||||||
|
}
|
||||||
|
|
||||||
|
delay, _ := fields["connecting_again_in"].(string)
|
||||||
|
got = append(got, delay)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Equal(got, delays) {
|
||||||
|
t.Errorf("logged failures to connect again in %v, want %v", got, delays)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// testCertificate returns a certificate for 127.0.0.1 that is its own
|
||||||
|
// CA, and a pool that holds it. It is valid on the bubble's clock, which
|
||||||
|
// starts at 2000-01-01T00:00:00Z.
|
||||||
|
func testCertificate(t *testing.T) (tls.Certificate, *x509.CertPool) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("generate a key: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
template := &x509.Certificate{
|
||||||
|
SerialNumber: big.NewInt(1),
|
||||||
|
Subject: pkix.Name{CommonName: "smallwebwaf test CA"},
|
||||||
|
NotBefore: time.Date(1999, 12, 31, 0, 0, 0, 0, time.UTC),
|
||||||
|
NotAfter: time.Date(2000, 1, 2, 0, 0, 0, 0, time.UTC),
|
||||||
|
IsCA: true,
|
||||||
|
BasicConstraintsValid: true,
|
||||||
|
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
|
||||||
|
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||||
|
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)},
|
||||||
|
}
|
||||||
|
|
||||||
|
der, err := x509.CreateCertificate(rand.Reader, template, template,
|
||||||
|
&key.PublicKey, key)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create a certificate: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
certificate, err := x509.ParseCertificate(der)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse the certificate: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
roots := x509.NewCertPool()
|
||||||
|
roots.AddCert(certificate)
|
||||||
|
|
||||||
|
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots
|
||||||
|
}
|
||||||
@@ -9,6 +9,8 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||||
)
|
)
|
||||||
|
|
||||||
// The action a request line names: what smallwebwaf did with the
|
// The action a request line names: what smallwebwaf did with the
|
||||||
@@ -24,8 +26,14 @@ const (
|
|||||||
// for, or whose answer could not be passed on.
|
// for, or whose answer could not be passed on.
|
||||||
ActionUpstreamError = "upstream_error"
|
ActionUpstreamError = "upstream_error"
|
||||||
// ActionRateLimited is a request refused because it took its client
|
// 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"
|
ActionRateLimited = "rate_limited"
|
||||||
|
// ActionBanned is a request refused because a ban covers its client,
|
||||||
|
// or because it matched a ban rule, which bans the client.
|
||||||
|
ActionBanned = "banned"
|
||||||
|
// ActionRuleBlocked is a request refused because it matched a block
|
||||||
|
// rule.
|
||||||
|
ActionRuleBlocked = "rule_blocked"
|
||||||
// ActionDenied is a request refused because its client is in
|
// ActionDenied is a request refused because its client is in
|
||||||
// SWWAF_DENY_NETS.
|
// SWWAF_DENY_NETS.
|
||||||
ActionDenied = "denied"
|
ActionDenied = "denied"
|
||||||
@@ -36,39 +44,101 @@ const (
|
|||||||
ActionAdmin = "admin"
|
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.
|
// timeLayout is RFC 3339 with milliseconds.
|
||||||
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
|
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
|
||||||
|
|
||||||
// Line is one request's line in the request log. The field names are
|
// Line is one request's line in the request log. The field names, and
|
||||||
// those of the "Request log" section of SPEC.md.
|
// their order, are those of the "Request log" section of SPEC.md. A field
|
||||||
|
// that may not apply to a request is left out of its line when it does
|
||||||
|
// not.
|
||||||
//
|
//
|
||||||
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
||||||
type Line struct {
|
type Line struct {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Time string `json:"time"`
|
|
||||||
ClientIP string `json:"client_ip"`
|
// The standard web log fields. Scheme is how the client reached
|
||||||
PeerIP string `json:"peer_ip"`
|
// smallwebwaf, or the trusted proxy in front of it.
|
||||||
Country string `json:"country"`
|
Time string `json:"time"`
|
||||||
Method string `json:"method"`
|
Instance string `json:"instance"`
|
||||||
Host string `json:"host"`
|
ClientIP string `json:"client_ip"`
|
||||||
Path string `json:"path"`
|
Method string `json:"method"`
|
||||||
Query string `json:"query"`
|
Scheme string `json:"scheme"`
|
||||||
Protocol string `json:"protocol"`
|
Host string `json:"host"`
|
||||||
Status int `json:"status"`
|
Path string `json:"path"`
|
||||||
UpstreamStatus int `json:"upstream_status,omitempty"`
|
Query string `json:"query"`
|
||||||
RequestBytes int64 `json:"request_bytes"`
|
Protocol string `json:"protocol"`
|
||||||
ResponseBytes int64 `json:"response_bytes"`
|
Status int `json:"status"`
|
||||||
Referer string `json:"referer"`
|
RequestBytes int64 `json:"request_bytes"`
|
||||||
UserAgent string `json:"user_agent"`
|
ResponseBytes int64 `json:"response_bytes"`
|
||||||
Action string `json:"action"`
|
Referer string `json:"referer"`
|
||||||
|
UserAgent string `json:"user_agent"`
|
||||||
|
|
||||||
|
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
|
||||||
|
// or a new one, and is sent on to the app. ForwardedFor is the
|
||||||
|
// X-Forwarded-For header as received. ClientGroup is the netblock the
|
||||||
|
// client is counted as.
|
||||||
|
RequestID string `json:"request_id"`
|
||||||
|
PeerIP string `json:"peer_ip"`
|
||||||
|
ForwardedFor string `json:"forwarded_for,omitempty"`
|
||||||
|
ClientGroup string `json:"client_group"`
|
||||||
|
Country string `json:"country"`
|
||||||
|
ContentType string `json:"content_type,omitempty"`
|
||||||
|
// ContentLength is the length of its body the request announced.
|
||||||
|
ContentLength int64 `json:"content_length,omitempty"`
|
||||||
|
// RequestHeaders are the headers SWWAF_LOG_REQUEST_HEADERS names that
|
||||||
|
// the request carried, by name in lower case.
|
||||||
|
RequestHeaders map[string]string `json:"request_headers,omitempty"`
|
||||||
|
HasAuthorization bool `json:"has_authorization,omitempty"`
|
||||||
|
HasCookie bool `json:"has_cookie,omitempty"`
|
||||||
|
// Websocket is true when the connection was upgraded, as for a
|
||||||
|
// WebSocket.
|
||||||
|
Websocket bool `json:"websocket,omitempty"`
|
||||||
|
|
||||||
|
// Response detail, from the headers of the answer: the app's, as
|
||||||
|
// passed on, or those of smallwebwaf's own. Aborted is true when the
|
||||||
|
// client went away early.
|
||||||
|
ResponseContentType string `json:"response_content_type,omitempty"`
|
||||||
|
UpstreamStatus int `json:"upstream_status,omitempty"`
|
||||||
|
CacheControl string `json:"cache_control,omitempty"`
|
||||||
|
Location string `json:"location,omitempty"`
|
||||||
|
Aborted bool `json:"aborted,omitempty"`
|
||||||
|
|
||||||
|
// The decision.
|
||||||
|
Action string `json:"action"`
|
||||||
|
// WouldAction is, in observe mode, the action enforce mode would have
|
||||||
|
// taken with a request it would have refused: ActionDenied,
|
||||||
|
// ActionBanned, ActionCountryDenied, ActionRateLimited or
|
||||||
|
// ActionRuleBlocked.
|
||||||
|
WouldAction string `json:"would_action,omitempty"`
|
||||||
|
// Counts are the client's requests as the rate limits counted them
|
||||||
|
// with this one, for a request they counted.
|
||||||
|
Counts ratelimit.Counts `json:"counts,omitzero"`
|
||||||
|
// RuleIDs are the ids of the rule file rules the request matched.
|
||||||
|
RuleIDs []string `json:"rule_ids,omitempty"`
|
||||||
// LimitHit is the window whose rate limit the request went over:
|
// LimitHit is the window whose rate limit the request went over:
|
||||||
// minute, hour or day.
|
// minute, hour or day.
|
||||||
LimitHit string `json:"limit_hit,omitempty"`
|
LimitHit string `json:"limit_hit,omitempty"`
|
||||||
// Aborted is true when the client went away early.
|
// Offence is the offence the request was held as, OffenceLimit.
|
||||||
Aborted bool `json:"aborted,omitempty"`
|
Offence string `json:"offence,omitempty"`
|
||||||
// DurationTotal and DurationUpstreamTotal are in milliseconds.
|
// BanExpires is when the ban the request made, or was refused under,
|
||||||
DurationTotal float64 `json:"duration_total"`
|
// ends: a time, or "permanent".
|
||||||
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"`
|
BanExpires string `json:"ban_expires,omitempty"`
|
||||||
|
|
||||||
|
// The timings, in milliseconds. DurationChecks is the time until the
|
||||||
|
// checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte
|
||||||
|
// and DurationUpstreamTotal run from when the request was handed to the
|
||||||
|
// app: until there was a connection to it, until the first byte of its
|
||||||
|
// answer arrived, and until the end. Each but DurationTotal is nil for
|
||||||
|
// a request that did not get that far.
|
||||||
|
DurationTotal float64 `json:"duration_total"`
|
||||||
|
DurationChecks *float64 `json:"duration_checks,omitempty"`
|
||||||
|
DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"`
|
||||||
|
DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"`
|
||||||
|
DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write writes line to w as one JSON line marked "type":"request".
|
// Write writes line to w as one JSON line marked "type":"request".
|
||||||
|
|||||||
@@ -50,7 +50,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
unset := []string{
|
unset := []string{
|
||||||
"upstream_status", "limit_hit", "aborted", "duration_upstream_total",
|
"forwarded_for", "content_type", "content_length", "request_headers",
|
||||||
|
"has_authorization", "has_cookie", "websocket", "response_content_type",
|
||||||
|
"upstream_status", "cache_control", "location", "aborted", "counts",
|
||||||
|
"limit_hit", "offence", "ban_expires", "duration_checks",
|
||||||
|
"duration_upstream_connect", "duration_upstream_first_byte",
|
||||||
|
"duration_upstream_total",
|
||||||
}
|
}
|
||||||
for _, name := range unset {
|
for _, name := range unset {
|
||||||
_, present := fields[name]
|
_, present := fields[name]
|
||||||
|
|||||||
@@ -0,0 +1,469 @@
|
|||||||
|
// Package rules reads the rule files: the plain text files in
|
||||||
|
// SWWAF_RULES_DIR, one rule to a line, that each request is checked
|
||||||
|
// against, as the "Rule files" section of SPEC.md describes. They are read
|
||||||
|
// at start, and again once the directory has had no change for a short
|
||||||
|
// time after one is edited, added or removed.
|
||||||
|
package rules
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/hex"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"regexp"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fsnotify/fsnotify"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The actions a rule takes when it matches.
|
||||||
|
const (
|
||||||
|
// ActionLog notes the match in the request log, and does nothing else.
|
||||||
|
ActionLog = "log"
|
||||||
|
// ActionBlock refuses the request with 403.
|
||||||
|
ActionBlock = "block"
|
||||||
|
// ActionBan refuses the request and bans the client's netblock: the
|
||||||
|
// request is a clear sign of attack.
|
||||||
|
ActionBan = "ban"
|
||||||
|
)
|
||||||
|
|
||||||
|
// extension ends the name of every rule file.
|
||||||
|
const extension = ".rules"
|
||||||
|
|
||||||
|
// quietTime is how long SWWAF_RULES_DIR must go without a change before
|
||||||
|
// the rule files are read again, so that a file still being written, such
|
||||||
|
// as one saved in place, appended to or copied in with scp, is read only
|
||||||
|
// once whole.
|
||||||
|
const quietTime = 2 * time.Second
|
||||||
|
|
||||||
|
// headerTarget starts the target that is one request header,
|
||||||
|
// header:<Name>.
|
||||||
|
const headerTarget = "header:"
|
||||||
|
|
||||||
|
// escapeLength is the length of a percent escape, such as %2e.
|
||||||
|
const escapeLength = 3
|
||||||
|
|
||||||
|
var (
|
||||||
|
// ruleLine is a rule: four fields separated by spaces or tabs, of
|
||||||
|
// which the fourth, the regex, runs to the end of the line.
|
||||||
|
ruleLine = regexp.MustCompile(`^([^ \t]+)[ \t]+([^ \t]+)[ \t]+([^ \t]+)[ \t]+(.+)$`)
|
||||||
|
// idChars are the characters of a rule's id.
|
||||||
|
idChars = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errNotRule = errors.New(
|
||||||
|
"is not a rule: an id, a target, an action and a regex, " +
|
||||||
|
"separated by spaces or tabs")
|
||||||
|
errNotID = errors.New("is not an id of letters, digits, - and _")
|
||||||
|
errNotTarget = errors.New(
|
||||||
|
"is not path, query, uri, method, host, user_agent, referer or header:<Name>")
|
||||||
|
errNotHeaderName = errors.New(
|
||||||
|
"has a character after header: that no header name can have")
|
||||||
|
errHeaderTakenOut = errors.New(
|
||||||
|
"names a header that Go's HTTP server takes out of every request, " +
|
||||||
|
"so a rule never sees it")
|
||||||
|
errNotAction = errors.New("is not log, block or ban")
|
||||||
|
errNotRegex = errors.New("does not compile")
|
||||||
|
errUsedTwice = errors.New("is already the id of the rule at")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Rule is one rule of a rule file.
|
||||||
|
type Rule struct {
|
||||||
|
// ID names the rule in the request log, the metrics and ban notes.
|
||||||
|
ID string
|
||||||
|
// Target is what the regex is matched against, such as path or
|
||||||
|
// header:Accept.
|
||||||
|
Target string
|
||||||
|
// Action is ActionLog, ActionBlock or ActionBan.
|
||||||
|
Action string
|
||||||
|
|
||||||
|
regex *regexp.Regexp
|
||||||
|
}
|
||||||
|
|
||||||
|
// Params are what Load needs.
|
||||||
|
type Params struct {
|
||||||
|
// Dir is the directory of the rule files (SWWAF_RULES_DIR).
|
||||||
|
Dir string
|
||||||
|
// Enabled is SWWAF_RULES_ENABLED: while it is false, no file is read
|
||||||
|
// and no rule loaded.
|
||||||
|
Enabled bool
|
||||||
|
// ProcessLog receives how many rules were read, and the error in a
|
||||||
|
// rule file edited while smallwebwaf runs.
|
||||||
|
ProcessLog *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// Files are the rule files of a running smallwebwaf, and the rules read
|
||||||
|
// from them. They are safe for concurrent use.
|
||||||
|
type Files struct {
|
||||||
|
params Params
|
||||||
|
// rules are the rules loaded, in the order of their files' names, and
|
||||||
|
// then of their lines.
|
||||||
|
rules atomic.Pointer[[]Rule]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load reads the rules of every *.rules file in Dir, in the order of the
|
||||||
|
// files' names, unless Enabled is false. A Dir that cannot be read is an
|
||||||
|
// error, and so is a line that is not a rule, a header name with a
|
||||||
|
// character no header name can have, a rule for the Host or the
|
||||||
|
// Transfer-Encoding header, which Go's HTTP server takes out of every
|
||||||
|
// request, a regex that does not compile and an id used twice, each named
|
||||||
|
// with its file and line.
|
||||||
|
func Load(params Params) (*Files, error) {
|
||||||
|
f := &Files{params: params}
|
||||||
|
f.rules.Store(&[]Rule{})
|
||||||
|
|
||||||
|
if !params.Enabled {
|
||||||
|
return f, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
rules, err := read(params.Dir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
f.rules.Store(&rules)
|
||||||
|
f.logRead(len(rules))
|
||||||
|
|
||||||
|
return f, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match checks r against the rules, in order, and returns those it
|
||||||
|
// matches, up to the first whose action refuses it, block or ban, which
|
||||||
|
// is then the last one returned.
|
||||||
|
func (f *Files) Match(r *http.Request) []Rule {
|
||||||
|
var matched []Rule
|
||||||
|
|
||||||
|
for _, rule := range *f.rules.Load() {
|
||||||
|
if !rule.matches(r) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
matched = append(matched, rule)
|
||||||
|
if rule.Action != ActionLog {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return matched
|
||||||
|
}
|
||||||
|
|
||||||
|
// Len returns how many rules are loaded.
|
||||||
|
func (f *Files) Len() int {
|
||||||
|
return len(*f.rules.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Watch watches Dir until ctx is done, and reads the rule files again
|
||||||
|
// once Dir has had no change for quietTime, after one is edited, added or
|
||||||
|
// removed, and after Watch starts watching. If they then hold an error,
|
||||||
|
// the rules stay as they were, the error is logged with its file and
|
||||||
|
// line, and the files are read again after the next change. If Dir cannot
|
||||||
|
// be watched, that is logged, and the rules stay as they were loaded.
|
||||||
|
// While Enabled is false, Watch returns at once.
|
||||||
|
func (f *Files) Watch(ctx context.Context) {
|
||||||
|
if !f.params.Enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
watcher, err := fsnotify.NewWatcher()
|
||||||
|
if err == nil {
|
||||||
|
defer func() {
|
||||||
|
_ = watcher.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
err = watcher.Add(f.params.Dir)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error("cannot watch the rule files for edits",
|
||||||
|
"error", err.Error())
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.ProcessLog.Info("watching the rule files for edits",
|
||||||
|
"directory", f.params.Dir)
|
||||||
|
|
||||||
|
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readAfterChanges reads the rule files again once quietTime has passed
|
||||||
|
// without a change from events, until ctx is done, and logs the errors
|
||||||
|
// from errs. The wait starts at once, as if for a change, so that an edit
|
||||||
|
// saved after Load read the files, and before Dir was watched, is taken
|
||||||
|
// in too.
|
||||||
|
func (f *Files) readAfterChanges(
|
||||||
|
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
|
||||||
|
) {
|
||||||
|
quiet := time.NewTimer(quietTime)
|
||||||
|
defer quiet.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-events:
|
||||||
|
quiet.Reset(quietTime)
|
||||||
|
case <-quiet.C:
|
||||||
|
f.readAgain()
|
||||||
|
case err := <-errs:
|
||||||
|
f.params.ProcessLog.Warn("watching the rule files failed",
|
||||||
|
"error", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// readAgain reads the rule files again, in place of the rules loaded, or
|
||||||
|
// logs the error that keeps the rules as they were.
|
||||||
|
func (f *Files) readAgain() {
|
||||||
|
rules, err := read(f.params.Dir)
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error(
|
||||||
|
"a rule file has an error, and the rules stay as they were",
|
||||||
|
"error", err.Error())
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f.rules.Store(&rules)
|
||||||
|
f.logRead(len(rules))
|
||||||
|
}
|
||||||
|
|
||||||
|
// logRead logs that the rule files were read, and how many rules they
|
||||||
|
// hold, which can be none.
|
||||||
|
func (f *Files) logRead(count int) {
|
||||||
|
f.params.ProcessLog.Info("read the rule files",
|
||||||
|
"directory", f.params.Dir, "rules", count)
|
||||||
|
}
|
||||||
|
|
||||||
|
// read returns the rules of every rule file in dir, in the order of the
|
||||||
|
// files' names, and then of their lines. A file whose name starts with a
|
||||||
|
// dot, such as an editor's lock file .#50-app.rules, is not a rule file,
|
||||||
|
// as a shell's *.rules would not match it.
|
||||||
|
func read(dir string) ([]Rule, error) {
|
||||||
|
entries, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var rules []Rule
|
||||||
|
|
||||||
|
// places are where each id is, as "<file>, line <n>".
|
||||||
|
places := map[string]string{}
|
||||||
|
|
||||||
|
for _, entry := range entries {
|
||||||
|
name := entry.Name()
|
||||||
|
if entry.IsDir() || strings.HasPrefix(name, ".") || filepath.Ext(name) != extension {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
rules, err = readFile(filepath.Join(dir, name), rules, places)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// readFile appends the rules of the rule file at path to rules. places
|
||||||
|
// are where each id read so far is, and gain those of the file.
|
||||||
|
func readFile(path string, rules []Rule, places map[string]string) ([]Rule, error) {
|
||||||
|
data, err := os.ReadFile(path) //nolint:gosec // a rule file, in SWWAF_RULES_DIR
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
number := 0
|
||||||
|
|
||||||
|
for line := range strings.Lines(string(data)) {
|
||||||
|
number++
|
||||||
|
place := fmt.Sprintf("%s, line %d", path, number)
|
||||||
|
|
||||||
|
text := strings.TrimSuffix(strings.TrimSuffix(line, "\n"), "\r")
|
||||||
|
|
||||||
|
rule, isRule, err := parse(text)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%s: %w", place, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isRule {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
first, used := places[rule.ID]
|
||||||
|
if used {
|
||||||
|
return nil, fmt.Errorf("%s: the id %q %w %s", place, rule.ID, errUsedTwice, first)
|
||||||
|
}
|
||||||
|
|
||||||
|
places[rule.ID] = place
|
||||||
|
rules = append(rules, rule)
|
||||||
|
}
|
||||||
|
|
||||||
|
return rules, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parse reads a line of a rule file. It returns false for a blank line
|
||||||
|
// and for a comment, a line that starts with #. Spaces and tabs at the
|
||||||
|
// end of the line are not part of its regex, so a line with only those
|
||||||
|
// after its action has no regex, and is not a rule.
|
||||||
|
func parse(line string) (Rule, bool, error) {
|
||||||
|
line = strings.Trim(line, " \t")
|
||||||
|
if line == "" || strings.HasPrefix(line, "#") {
|
||||||
|
return Rule{}, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fields := ruleLine.FindStringSubmatch(line)
|
||||||
|
if fields == nil {
|
||||||
|
return Rule{}, false, errNotRule
|
||||||
|
}
|
||||||
|
|
||||||
|
rule := Rule{ID: fields[1], Target: fields[2], Action: fields[3]}
|
||||||
|
headerName, isHeader := strings.CutPrefix(rule.Target, headerTarget)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case !idChars.MatchString(rule.ID):
|
||||||
|
return Rule{}, false, fmt.Errorf("the id %q %w", rule.ID, errNotID)
|
||||||
|
case !isTarget(rule.Target):
|
||||||
|
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotTarget)
|
||||||
|
case isHeader && !config.IsHeaderName(headerName):
|
||||||
|
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotHeaderName)
|
||||||
|
case strings.EqualFold(rule.Target, headerTarget+"Host"):
|
||||||
|
return Rule{}, false, fmt.Errorf(
|
||||||
|
"the target %q %w; the request's host is the target host",
|
||||||
|
rule.Target, errHeaderTakenOut)
|
||||||
|
case strings.EqualFold(rule.Target, headerTarget+"Transfer-Encoding"):
|
||||||
|
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errHeaderTakenOut)
|
||||||
|
case !slices.Contains([]string{ActionLog, ActionBlock, ActionBan}, rule.Action):
|
||||||
|
return Rule{}, false, fmt.Errorf("the action %q %w", rule.Action, errNotAction)
|
||||||
|
}
|
||||||
|
|
||||||
|
regex, err := regexp.Compile(fields[4])
|
||||||
|
if err != nil {
|
||||||
|
return Rule{}, false, fmt.Errorf("the regex %w: %w", errNotRegex, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rule.regex = regex
|
||||||
|
|
||||||
|
return rule, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// isTarget reports whether target is one a rule may have.
|
||||||
|
func isTarget(target string) bool {
|
||||||
|
switch target {
|
||||||
|
case "path", "query", "uri", "method", "host", "user_agent", "referer":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
name, isHeader := strings.CutPrefix(target, headerTarget)
|
||||||
|
|
||||||
|
return isHeader && name != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// matches reports whether the rule's regex matches its target in r. For
|
||||||
|
// uri it is matched against the path and query as received, and against
|
||||||
|
// them once percent-decoded, so that an encoded probe cannot slip past.
|
||||||
|
func (rule Rule) matches(r *http.Request) bool {
|
||||||
|
if rule.Target == "uri" {
|
||||||
|
uri := pathAndQuery(r)
|
||||||
|
|
||||||
|
return rule.regex.MatchString(uri) || rule.regex.MatchString(decodeOnce(uri))
|
||||||
|
}
|
||||||
|
|
||||||
|
return rule.regex.MatchString(value(rule.Target, r))
|
||||||
|
}
|
||||||
|
|
||||||
|
// value returns what a rule with target, other than uri, is matched
|
||||||
|
// against in r: the path and the query as the client sent them, before
|
||||||
|
// any decoding or re-encoding, split at the first ?, and a header's values
|
||||||
|
// joined by ", ", as HTTP joins those of a header sent more than once.
|
||||||
|
func value(target string, r *http.Request) string {
|
||||||
|
switch target {
|
||||||
|
case "path":
|
||||||
|
path, _, _ := strings.Cut(pathAndQuery(r), "?")
|
||||||
|
|
||||||
|
return path
|
||||||
|
case "query":
|
||||||
|
_, query, _ := strings.Cut(pathAndQuery(r), "?")
|
||||||
|
|
||||||
|
return query
|
||||||
|
case "method":
|
||||||
|
return r.Method
|
||||||
|
case "host":
|
||||||
|
return r.Host
|
||||||
|
case "user_agent":
|
||||||
|
return header(r, "User-Agent")
|
||||||
|
case "referer":
|
||||||
|
return header(r, "Referer")
|
||||||
|
default:
|
||||||
|
return header(r, strings.TrimPrefix(target, headerTarget))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// pathAndQuery returns the target of r's request line, r.RequestURI, as
|
||||||
|
// the client sent it, less any scheme and host: a target with a scheme
|
||||||
|
// gives what follows the scheme and its :, and the host when // follows.
|
||||||
|
// So http://host/path, as a client sends it to a proxy, gives /path, and
|
||||||
|
// so does http:/path, which Go reads as a target with a scheme and no
|
||||||
|
// host. r.URL is not used: when the path holds a character it escapes,
|
||||||
|
// such as \ or a non-ASCII byte, it decodes the whole path and escapes it
|
||||||
|
// again, so that \ becomes %5C and %2e a dot.
|
||||||
|
func pathAndQuery(r *http.Request) string {
|
||||||
|
if !r.URL.IsAbs() {
|
||||||
|
return r.RequestURI
|
||||||
|
}
|
||||||
|
|
||||||
|
_, afterScheme, _ := strings.Cut(r.RequestURI, ":")
|
||||||
|
|
||||||
|
hostAndRest, hasHost := strings.CutPrefix(afterScheme, "//")
|
||||||
|
if !hasHost {
|
||||||
|
return afterScheme
|
||||||
|
}
|
||||||
|
|
||||||
|
start := strings.IndexAny(hostAndRest, "/?")
|
||||||
|
if start < 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return hostAndRest[start:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// header returns the values of r's header name joined by ", ", or "" if
|
||||||
|
// r has no such header.
|
||||||
|
func header(r *http.Request, name string) string {
|
||||||
|
return strings.Join(r.Header.Values(name), ", ")
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeOnce returns s with each percent escape, such as %2e, replaced by
|
||||||
|
// the byte it stands for. A % that is not followed by two hex digits is
|
||||||
|
// left as it is, so that a malformed escape cannot keep the rest of s
|
||||||
|
// from being decoded.
|
||||||
|
func decodeOnce(s string) string {
|
||||||
|
var decoded strings.Builder
|
||||||
|
|
||||||
|
for i := 0; i < len(s); i++ {
|
||||||
|
if s[i] == '%' && i+escapeLength <= len(s) {
|
||||||
|
b, err := hex.DecodeString(s[i+1 : i+escapeLength])
|
||||||
|
if err == nil {
|
||||||
|
decoded.Write(b)
|
||||||
|
|
||||||
|
i += escapeLength - 1
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded.WriteByte(s[i])
|
||||||
|
}
|
||||||
|
|
||||||
|
return decoded.String()
|
||||||
|
}
|
||||||
@@ -0,0 +1,662 @@
|
|||||||
|
package rules_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"log/slog"
|
||||||
|
"maps"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// What the process log says once Watch watches the directory, after
|
||||||
|
// each reading of the rule files, and for one that has an error.
|
||||||
|
watching = "watching the rule files for edits"
|
||||||
|
read = "read the rule files"
|
||||||
|
hasError = "a rule file has an error, and the rules stay as they were"
|
||||||
|
// maxLogLines is how many lines of the process log wait for a test to
|
||||||
|
// read them.
|
||||||
|
maxLogLines = 64
|
||||||
|
// browser is the user agent of an ordinary visitor.
|
||||||
|
browser = "Mozilla/5.0 (X11; Linux x86_64; rv:140.0) Gecko/20100101 Firefox/140.0"
|
||||||
|
// testFile is the rule file of a test that needs only one, and
|
||||||
|
// firstFile the first of a test's rule files.
|
||||||
|
testFile = "test.rules"
|
||||||
|
firstFile = "00-a.rules"
|
||||||
|
// userAgent is the header that carries the user agent.
|
||||||
|
userAgent = "User-Agent"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEachTargetMatchesWhatItNames(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
rule string // its target, action and regex
|
||||||
|
uri string // the request's path and query
|
||||||
|
header http.Header
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"path as received", `path log ^/%2eenv$`, "/%2eenv", nil, true},
|
||||||
|
{"path not decoded", `path log ^/\.env$`, "/%2eenv", nil, false},
|
||||||
|
{"path without the query", `path log ^/a$`, "/a?b=c", nil, true},
|
||||||
|
{"query as received", `query log ^b=%2e$`, "/a?b=%2e", nil, true},
|
||||||
|
{"uri as received", `uri log ^/a\?b=%2e$`, "/a?b=%2e", nil, true},
|
||||||
|
{"uri decoded", `uri log (\.\./){2}`, "/a?f=%2e%2e%2f%2e%2e%2f", nil, true},
|
||||||
|
{
|
||||||
|
"uri decoded past malformed escapes", `uri log (\.\./){2}&h=%$`,
|
||||||
|
"/a?g=%zz&f=%2e%2e%2f%2e%2e%2f&h=%", nil, true,
|
||||||
|
},
|
||||||
|
{"uri decoded only once", `uri log ^/a\.b$`, "/a%252eb", nil, false},
|
||||||
|
{"method", `method log ^PUT$`, "/", nil, true},
|
||||||
|
{"host", `host log ^app\.example$`, "/", nil, true},
|
||||||
|
{
|
||||||
|
"user_agent", `user_agent log ^sqlmap/`, "/",
|
||||||
|
http.Header{userAgent: {"sqlmap/1.8"}}, true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"user_agent sent twice", `user_agent log ^curl/8, sqlmap/`, "/",
|
||||||
|
http.Header{userAgent: {"curl/8", "sqlmap/1.8"}}, true,
|
||||||
|
},
|
||||||
|
{"user_agent missing", `user_agent log ^$`, "/", nil, true},
|
||||||
|
{
|
||||||
|
"referer", `referer log ^https://spam\.example/`, "/",
|
||||||
|
http.Header{"Referer": {"https://spam.example/buy"}}, true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a header sent twice", `header:x-api-version log ^2, 3$`, "/",
|
||||||
|
http.Header{"X-Api-Version": {"2", "3"}}, true,
|
||||||
|
},
|
||||||
|
{"a header missing", `header:X-Api-Version log ^$`, "/", nil, true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
files := load(t, ruleFiles{testFile: "a-rule " + tc.rule + "\n"})
|
||||||
|
|
||||||
|
// Every request is a PUT, which the method rule looks for.
|
||||||
|
r := httptest.NewRequestWithContext(t.Context(), http.MethodPut,
|
||||||
|
"http://app.example"+tc.uri, nil)
|
||||||
|
maps.Copy(r.Header, tc.header)
|
||||||
|
|
||||||
|
got := len(files.Match(r)) == 1
|
||||||
|
if got != tc.want {
|
||||||
|
t.Errorf("%s matches %s: %t, want %t", tc.rule, tc.uri, got, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPathMatchedAsTheClientSentIt(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Each path holds a character Go's URL type would escape again, \ or
|
||||||
|
// a non-ASCII byte, and each rule is written for the path as sent.
|
||||||
|
for _, tc := range []struct {
|
||||||
|
rule string // its target, action and regex
|
||||||
|
sent string // the path and query the client sent
|
||||||
|
}{
|
||||||
|
{`path log ^/\.\.\\\.\.\\windows\\win\.ini$`, `/..\..\windows\win.ini`},
|
||||||
|
{`path log ^/%2e%2e\\%2e%2e\\windows\\win\.ini$`, `/%2e%2e\%2e%2e\windows\win.ini`},
|
||||||
|
{`path log ^/café$`, "/café?x=1"},
|
||||||
|
{`uri log ^/%2e%2e\\%2e%2e\\boot\.ini\?x=1$`, `/%2e%2e\%2e%2e\boot.ini?x=1`},
|
||||||
|
} {
|
||||||
|
files := load(t, ruleFiles{testFile: "as-sent " + tc.rule + "\n"})
|
||||||
|
|
||||||
|
// The target in origin form, as traefik sends it, in absolute form,
|
||||||
|
// as a client sends it to a proxy, and with a scheme but no host,
|
||||||
|
// which Go reads as absolute form with no host, sending the app
|
||||||
|
// the path.
|
||||||
|
for _, target := range []string{
|
||||||
|
tc.sent, "http://app.example" + tc.sent, "http:" + tc.sent, "foo:" + tc.sent,
|
||||||
|
} {
|
||||||
|
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet, target, nil)
|
||||||
|
wantMatched(t, files, r, "as-sent")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchingStopsAtTheFirstRuleThatRefuses(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
files := load(t, ruleFiles{testFile: `
|
||||||
|
every-path path log ^/
|
||||||
|
no-path path log ^$
|
||||||
|
first-refusal path block ^/probe
|
||||||
|
later-ban path ban ^/probe
|
||||||
|
after path log ^/
|
||||||
|
`})
|
||||||
|
|
||||||
|
// Every log rule that matches is noted, and the block rule ends the
|
||||||
|
// matching.
|
||||||
|
wantMatched(t, files, get(t, "/probe"), "every-path", "first-refusal")
|
||||||
|
wantMatched(t, files, get(t, "/page"), "every-path", "after")
|
||||||
|
|
||||||
|
// A ban rule ends it too.
|
||||||
|
files = load(t, ruleFiles{testFile: "ban path ban ^/\nlater path block ^/\n"})
|
||||||
|
wantMatched(t, files, get(t, "/"), "ban")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSpacesAndTabsEndingALineAreNotPartOfItsRegex(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
files := load(t, ruleFiles{testFile: "env-file path block ^/\\.env$ \t \n"})
|
||||||
|
|
||||||
|
wantMatched(t, files, get(t, "/.env"), "env-file")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilesReadInNameOrderThenLineOrder(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
files := load(t, ruleFiles{
|
||||||
|
"50-b.rules": "b1 path log ^/\n\n# a comment\n # an indented one\nb2 path log ^/\n",
|
||||||
|
firstFile: "a1 path log ^/\r\n",
|
||||||
|
// None is a rule file.
|
||||||
|
"notes.txt": "notes, not rules\n",
|
||||||
|
"10-c.rules.bak": "an old copy\n",
|
||||||
|
"20-d.rules/keep": "a file in a directory\n",
|
||||||
|
})
|
||||||
|
|
||||||
|
wantMatched(t, files, get(t, "/"), "a1", "b1", "b2")
|
||||||
|
|
||||||
|
if files.Len() != 3 {
|
||||||
|
t.Errorf("%d rules loaded, want 3", files.Len())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFileWhoseNameStartsWithADotIsNotARuleFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := writeFiles(t, ruleFiles{firstFile: "probe path block ^/probe\n"})
|
||||||
|
|
||||||
|
// The lock file Emacs makes beside a file while it is edited: a link to
|
||||||
|
// nothing, which cannot be read.
|
||||||
|
err := os.Symlink("user@host.1234:1700000000", filepath.Join(dir, ".#"+firstFile))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("symlink: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
params, _ := newParams(dir)
|
||||||
|
|
||||||
|
files, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantMatched(t, files, get(t, "/probe"), "probe")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFaultStopsTheStartNamingTheFileAndLine(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
content string
|
||||||
|
line int
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
"too few fields", "env-file path ban\n", 1,
|
||||||
|
"is not a rule: an id, a target, an action and a regex, " +
|
||||||
|
"separated by spaces or tabs",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Else its regex would be a space, found in nearly every user agent.
|
||||||
|
"a regex of only spaces and tabs", "scanner user_agent ban\t \n", 1,
|
||||||
|
"is not a rule: an id, a target, an action and a regex, " +
|
||||||
|
"separated by spaces or tabs",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an id of other characters", "# ids\n\nenv.file path ban ^/\n", 3,
|
||||||
|
`the id "env.file" is not an id of letters, digits, - and _`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an unknown target", "env-file paths ban ^/\n", 1,
|
||||||
|
`the target "paths" is not path, query, uri, method, host, ` +
|
||||||
|
"user_agent, referer or header:<Name>",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a header without a name", "env-file header: ban ^/\n", 1,
|
||||||
|
`the target "header:" is not path, query, uri, method, host, ` +
|
||||||
|
"user_agent, referer or header:<Name>",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a header name written with its colon", "sqlmap header:User-Agent: ban sqlmap\n", 1,
|
||||||
|
`the target "header:User-Agent:" has a character after header: ` +
|
||||||
|
"that no header name can have",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a header name with a semicolon", "accept header:Accept;q log ^$\n", 1,
|
||||||
|
`the target "header:Accept;q" has a character after header: ` +
|
||||||
|
"that no header name can have",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a header name with brackets", "x-header header:X(y) log ^$\n", 1,
|
||||||
|
`the target "header:X(y)" has a character after header: ` +
|
||||||
|
"that no header name can have",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"the Host header", "host-header header:host block ^$\n", 1,
|
||||||
|
`the target "header:host" names a header that Go's HTTP server ` +
|
||||||
|
"takes out of every request, so a rule never sees it; " +
|
||||||
|
"the request's host is the target host",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"the Transfer-Encoding header",
|
||||||
|
"# bodies sent in chunks\nchunked header:Transfer-Encoding block ^chunked$\n", 2,
|
||||||
|
`the target "header:Transfer-Encoding" names a header that Go's ` +
|
||||||
|
"HTTP server takes out of every request, so a rule never sees it",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"an unknown action", "env-file path deny ^/\n", 1,
|
||||||
|
`the action "deny" is not log, block or ban`,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"a regex that does not compile", "env-file path ban ^/(\n", 1,
|
||||||
|
"the regex does not compile: error parsing regexp: " +
|
||||||
|
"missing closing ): `^/(`",
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := writeFiles(t, ruleFiles{"00-default.rules": tc.content})
|
||||||
|
path := filepath.Join(dir, "00-default.rules")
|
||||||
|
|
||||||
|
wantRefused(t, dir, path+", line "+strconv.Itoa(tc.line)+": "+tc.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIDUsedTwiceStopsTheStartNamingBothPlaces(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := writeFiles(t, ruleFiles{
|
||||||
|
"00-a.rules": "probe path log ^/a\n",
|
||||||
|
"50-b.rules": "other path log ^/b\nprobe path ban ^/c\n",
|
||||||
|
})
|
||||||
|
|
||||||
|
wantRefused(t, dir, filepath.Join(dir, "50-b.rules")+`, line 2: the id "probe" `+
|
||||||
|
"is already the id of the rule at "+filepath.Join(dir, "00-a.rules")+", line 1")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDirectoryThatDoesNotExistStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := filepath.Join(t.TempDir(), "rules.d")
|
||||||
|
|
||||||
|
wantRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
|
||||||
|
": no such file or directory")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEmptyDirectoryLoadsNoRulesAndSaysSo(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
params, lines := newParams(writeFiles(t, ruleFiles{"00-default.rules": "# none\n"}))
|
||||||
|
|
||||||
|
files, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
line := lines.waitFor(t, read)
|
||||||
|
if files.Len() != 0 || line["rules"] != 0.0 {
|
||||||
|
t.Errorf("%d rules loaded, and the log says %v, want none", files.Len(), line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRuleFilesOffReadNothing(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// SWWAF_RULES_DIR does not exist, which would stop the start.
|
||||||
|
params, _ := newParams(filepath.Join(t.TempDir(), "rules.d"))
|
||||||
|
params.Enabled = false
|
||||||
|
|
||||||
|
files, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if files.Len() != 0 || files.Match(get(t, "/")) != nil {
|
||||||
|
t.Errorf("%d rules loaded with the rule files off", files.Len())
|
||||||
|
}
|
||||||
|
|
||||||
|
// It would watch until the test ends.
|
||||||
|
files.Watch(t.Context())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEditsTakenInWhileRunning(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
||||||
|
files, lines := watch(t, dir)
|
||||||
|
|
||||||
|
// matches reports whether path matches a rule.
|
||||||
|
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
|
||||||
|
|
||||||
|
// A file added.
|
||||||
|
save(t, dir, "50-b.rules", "second path block ^/second\n")
|
||||||
|
lines.waitUntil(t, func() bool { return matches("/second") })
|
||||||
|
wantMatched(t, files, get(t, "/first"), "first")
|
||||||
|
|
||||||
|
// A file edited.
|
||||||
|
save(t, dir, firstFile, "first path block ^/edited\n")
|
||||||
|
lines.waitUntil(t, func() bool { return !matches("/first") })
|
||||||
|
wantMatched(t, files, get(t, "/edited"), "first")
|
||||||
|
|
||||||
|
// A file removed.
|
||||||
|
err := os.Remove(filepath.Join(dir, "50-b.rules"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("remove: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
lines.waitUntil(t, func() bool { return !matches("/second") })
|
||||||
|
wantMatched(t, files, get(t, "/edited"), "first")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
||||||
|
files, lines := watch(t, dir)
|
||||||
|
|
||||||
|
// The edit's second line has an unknown action, so the rules stay as
|
||||||
|
// they were, the first line's earlier version included.
|
||||||
|
save(t, dir, firstFile, "first path block ^/edited\nsecond path bann ^/second\n")
|
||||||
|
|
||||||
|
line := lines.waitFor(t, hasError)
|
||||||
|
want := filepath.Join(dir, firstFile) +
|
||||||
|
`, line 2: the action "bann" is not log, block or ban`
|
||||||
|
|
||||||
|
if line["error"] != want || line["level"] != "ERROR" {
|
||||||
|
t.Errorf("logged %v, want an error %q", line, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantMatched(t, files, get(t, "/first"), "first")
|
||||||
|
wantMatched(t, files, get(t, "/second"))
|
||||||
|
|
||||||
|
// Once mended, the file is read again.
|
||||||
|
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
|
||||||
|
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
|
||||||
|
wantMatched(t, files, get(t, "/edited"), "first")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
params, _ := newParams(filepath.Join("..", "..", "share", "rules.d"))
|
||||||
|
|
||||||
|
files, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load the default file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Probes sent by a browser, by the rule that refuses them.
|
||||||
|
for rule, targets := range map[string][]string{
|
||||||
|
"env-file": {"/.env", "/.env.production", "/.ENV"},
|
||||||
|
"vcs-dir": {"/.git/config", "/.git", "/.svn/entries"},
|
||||||
|
"secrets-dir": {"/.aws/credentials", "/.ssh/id_rsa"},
|
||||||
|
"secret-file": {"/.htpasswd", "/.DS_Store", "/.git-credentials"},
|
||||||
|
"editor-dir": {"/.vscode/sftp.json"},
|
||||||
|
"backup-file": {
|
||||||
|
"/wp-config.php.bak", "/index.php~", "/dump.sql", "/backup.sql.gz",
|
||||||
|
},
|
||||||
|
"log-file": {"/debug.log"},
|
||||||
|
"compose-file": {"/docker-compose.yml", "/compose.yaml"},
|
||||||
|
"php-shell": {"/shell.php"},
|
||||||
|
"path-traversal": {
|
||||||
|
"/static/../../etc/passwd", "/f?f=%2e%2e%2f%2e%2e%2fetc%2fpasswd",
|
||||||
|
},
|
||||||
|
} {
|
||||||
|
for _, target := range targets {
|
||||||
|
wantRefusedBy(t, files, target, browser, rule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Scanners, by their user agents.
|
||||||
|
for _, scanner := range []string{
|
||||||
|
"sqlmap/1.8.4#stable (https://sqlmap.org)",
|
||||||
|
"Mozilla/5.0 (compatible; Nuclei - Open-source project)",
|
||||||
|
} {
|
||||||
|
wantRefusedBy(t, files, "/", scanner, "scanner-agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ordinary requests to a code forge for files of those names deeper
|
||||||
|
// in its paths, and for other files at its root.
|
||||||
|
for _, target := range []string{
|
||||||
|
"/owner/repo/src/branch/main/.env.example",
|
||||||
|
"/owner/repo/src/branch/main/.env",
|
||||||
|
"/owner/repo/src/branch/main/.github/workflows/ci.yml",
|
||||||
|
"/owner/repo/src/branch/main/.vscode/settings.json",
|
||||||
|
"/owner/repo/src/branch/main/.htaccess",
|
||||||
|
"/owner/repo/src/branch/main/docker-compose.yml",
|
||||||
|
"/owner/repo/src/branch/main/db/schema.sql",
|
||||||
|
"/owner/repo/raw/branch/main/debug.log",
|
||||||
|
"/owner/repo.git/info/refs?service=git-upload-pack",
|
||||||
|
"/owner/repo/src/branch/main/docs/../README.md",
|
||||||
|
"/user/login?redirect_to=%2fowner%2frepo",
|
||||||
|
"/index.php",
|
||||||
|
"/.well-known/security.txt",
|
||||||
|
} {
|
||||||
|
r := get(t, target)
|
||||||
|
r.Header.Set(userAgent, browser)
|
||||||
|
|
||||||
|
matched := files.Match(r)
|
||||||
|
if len(matched) != 0 {
|
||||||
|
t.Errorf("%s matched %v, want no rule", target, ids(matched))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A request without a user agent is only noted.
|
||||||
|
wantMatched(t, files, get(t, "/"), "empty-agent")
|
||||||
|
}
|
||||||
|
|
||||||
|
// ruleFiles are files to write into a directory of rule files, by name.
|
||||||
|
type ruleFiles map[string]string
|
||||||
|
|
||||||
|
// writeFiles writes files into a new directory, and returns it.
|
||||||
|
func writeFiles(t *testing.T, files ruleFiles) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
for name, content := range files {
|
||||||
|
path := filepath.Join(dir, name)
|
||||||
|
|
||||||
|
err := os.MkdirAll(filepath.Dir(path), 0o700)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("mkdir: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = os.WriteFile(path, []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return dir
|
||||||
|
}
|
||||||
|
|
||||||
|
// save writes content to the rule file name in dir as an editor that
|
||||||
|
// saves by renaming does, so that the file is never seen half written.
|
||||||
|
func save(t *testing.T, dir, name, content string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
path := filepath.Join(dir, name)
|
||||||
|
|
||||||
|
err := os.WriteFile(path+".tmp", []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = os.Rename(path+".tmp", path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("rename: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newParams returns Params for the rule files in dir, switched on, with
|
||||||
|
// the process log in the processLog returned.
|
||||||
|
func newParams(dir string) (rules.Params, processLog) {
|
||||||
|
lines := make(processLog, maxLogLines)
|
||||||
|
|
||||||
|
return rules.Params{
|
||||||
|
Dir: dir,
|
||||||
|
Enabled: true,
|
||||||
|
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
|
||||||
|
}, lines
|
||||||
|
}
|
||||||
|
|
||||||
|
// load writes files into a new directory and loads the rules in it.
|
||||||
|
func load(t *testing.T, files ruleFiles) *rules.Files {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
params, _ := newParams(writeFiles(t, files))
|
||||||
|
params.ProcessLog = slog.New(slog.DiscardHandler)
|
||||||
|
|
||||||
|
loaded, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return loaded
|
||||||
|
}
|
||||||
|
|
||||||
|
// watch loads the rules in dir, runs their Watch until the test ends, and
|
||||||
|
// waits until it watches the directory.
|
||||||
|
func watch(t *testing.T, dir string) (*rules.Files, processLog) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
params, lines := newParams(dir)
|
||||||
|
|
||||||
|
files, err := rules.Load(params)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.Watch(ctx)
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
stop()
|
||||||
|
<-stopped
|
||||||
|
})
|
||||||
|
|
||||||
|
lines.waitFor(t, watching)
|
||||||
|
|
||||||
|
return files, lines
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRefused checks that loading the rule files in dir fails with the
|
||||||
|
// error want.
|
||||||
|
func wantRefused(t *testing.T, dir, want string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
params, _ := newParams(dir)
|
||||||
|
|
||||||
|
_, err := rules.Load(params)
|
||||||
|
if err == nil || err.Error() != want {
|
||||||
|
t.Errorf("error %v, want %s", err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// get returns a GET request for target, a path and an optional query, as
|
||||||
|
// smallwebwaf's server reads it, without a user agent.
|
||||||
|
func get(t *testing.T, target string) *http.Request {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
return httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"http://app.example"+target, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRefusedBy checks that a GET request for target with the user agent
|
||||||
|
// sent matches rule alone, and that rule refuses it.
|
||||||
|
func wantRefusedBy(t *testing.T, files *rules.Files, target, sent, rule string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
r := get(t, target)
|
||||||
|
r.Header.Set(userAgent, sent)
|
||||||
|
|
||||||
|
matched := files.Match(r)
|
||||||
|
if len(matched) != 1 || matched[0].ID != rule || matched[0].Action == rules.ActionLog {
|
||||||
|
t.Errorf("%s from %q matched %v, want %s alone, refusing it", target,
|
||||||
|
sent, ids(matched), rule)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantMatched checks the ids of the rules r matches, in order.
|
||||||
|
func wantMatched(t *testing.T, files *rules.Files, r *http.Request, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
got := ids(files.Match(r))
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("%s matched %v, want %v", r.URL, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ids returns the ids of matched.
|
||||||
|
func ids(matched []rules.Rule) []string {
|
||||||
|
got := make([]string, 0, len(matched))
|
||||||
|
for _, rule := range matched {
|
||||||
|
got = append(got, rule.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
return got
|
||||||
|
}
|
||||||
|
|
||||||
|
// processLog receives the lines of a process log, each a JSON object, for
|
||||||
|
// a test to wait for.
|
||||||
|
type processLog chan string
|
||||||
|
|
||||||
|
// Write receives a line of the process log.
|
||||||
|
func (l processLog) Write(line []byte) (int, error) {
|
||||||
|
l <- string(line)
|
||||||
|
|
||||||
|
return len(line), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitFor returns the next line of the process log whose message is msg,
|
||||||
|
// passing over the lines before it. It waits as long as that takes, so
|
||||||
|
// that a slow test process cannot fail the test.
|
||||||
|
func (l processLog) waitFor(t *testing.T, msg string) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for line := range l {
|
||||||
|
var fields map[string]any
|
||||||
|
|
||||||
|
err := json.Unmarshal([]byte(line), &fields)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("process log line %q is not JSON: %v", line, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if fields["msg"] == msg {
|
||||||
|
return fields
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitUntil waits for the rule files to be read until done reports true,
|
||||||
|
// as it does once they have been read after the test's last change. They
|
||||||
|
// can be read before then too, as they are once Watch starts watching.
|
||||||
|
func (l processLog) waitUntil(t *testing.T, done func() bool) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for !done() {
|
||||||
|
l.waitFor(t, read)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
package rules
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fsnotify/fsnotify"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The tests below run readAfterChanges 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 readAfterChanges waits again, so that every
|
||||||
|
// reading due by then is done. The test sends the changes itself, as the
|
||||||
|
// watch of a directory cannot run in a bubble.
|
||||||
|
|
||||||
|
func TestFileWrittenInTwoPartsTakenInOnlyWhole(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "50-app.rules")
|
||||||
|
writeFile(t, path, "first path block ^/first\n")
|
||||||
|
files := load(t, dir)
|
||||||
|
changes := run(t, files)
|
||||||
|
|
||||||
|
file, err := os.Create(path) //nolint:gosec // a file the test wrote
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = file.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
// The first part ends in the middle of a ban rule's regex, which,
|
||||||
|
// read then, would ban every request.
|
||||||
|
write(t, file, "first path block ^/first\nprobe path ban ^/")
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
time.Sleep(quietTime - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantMatched(t, files, "/anything")
|
||||||
|
|
||||||
|
// The second part starts the wait again.
|
||||||
|
write(t, file, `\.env$`+"\n")
|
||||||
|
|
||||||
|
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||||
|
|
||||||
|
time.Sleep(quietTime - time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantMatched(t, files, "/.env")
|
||||||
|
|
||||||
|
time.Sleep(time.Nanosecond)
|
||||||
|
synctest.Wait()
|
||||||
|
wantMatched(t, files, "/.env", "probe")
|
||||||
|
wantMatched(t, files, "/anything")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEditSavedBeforeTheWatchStartsTakenIn(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "50-app.rules")
|
||||||
|
writeFile(t, path, "first path block ^/first\n")
|
||||||
|
files := load(t, dir)
|
||||||
|
|
||||||
|
// Saved after Load read the files, and before the directory was
|
||||||
|
// watched, so that no change is seen for it.
|
||||||
|
writeFile(t, path, "first path block ^/edited\n")
|
||||||
|
run(t, files)
|
||||||
|
time.Sleep(quietTime)
|
||||||
|
synctest.Wait()
|
||||||
|
wantMatched(t, files, "/edited", "first")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// load loads the rules in dir.
|
||||||
|
func load(t *testing.T, dir string) *Files {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
files, err := Load(Params{
|
||||||
|
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("load: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return files
|
||||||
|
}
|
||||||
|
|
||||||
|
// run runs files' readAfterChanges until the test ends, and returns the
|
||||||
|
// channel that sends it changes.
|
||||||
|
func run(t *testing.T, files *Files) chan<- fsnotify.Event {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
changes := make(chan fsnotify.Event)
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
stopped := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.readAfterChanges(ctx, changes, nil)
|
||||||
|
close(stopped)
|
||||||
|
}()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
stop()
|
||||||
|
<-stopped
|
||||||
|
})
|
||||||
|
|
||||||
|
return changes
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeFile writes content to the file at path.
|
||||||
|
func writeFile(t *testing.T, path, content string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// write writes text to the end of file.
|
||||||
|
func write(t *testing.T, file *os.File, text string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
_, err := file.WriteString(text)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantMatched checks the ids of the rules that a GET request for path
|
||||||
|
// matches, in order.
|
||||||
|
func wantMatched(t *testing.T, files *Files, path string, want ...string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"http://app.example"+path, nil)
|
||||||
|
|
||||||
|
matched := files.Match(r)
|
||||||
|
|
||||||
|
got := make([]string, 0, len(matched))
|
||||||
|
for _, rule := range matched {
|
||||||
|
got = append(got, rule.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("%s matched %v, want %v", path, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -23,6 +23,8 @@ var errHealthEndpoint = errors.New("smallwebwaf's health endpoint answered")
|
|||||||
// smallwebwaf answers its health endpoint on 127.0.0.1, at the port in
|
// smallwebwaf answers its health endpoint on 127.0.0.1, at the port in
|
||||||
// SWWAF_LISTEN_ADDR, and the app accepts connections at the address in
|
// SWWAF_LISTEN_ADDR, and the app accepts connections at the address in
|
||||||
// SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1.
|
// SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1.
|
||||||
|
// It reads no other setting, nor a file that another names, so neither
|
||||||
|
// can fail it.
|
||||||
// args are the arguments after `healthcheck`; it takes none, and given
|
// args are the arguments after `healthcheck`; it takes none, and given
|
||||||
// one it names it on stderr and returns 1 without checking anything.
|
// one it names it on stderr and returns 1 without checking anything.
|
||||||
func HealthCheck(
|
func HealthCheck(
|
||||||
@@ -50,13 +52,13 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
|
|||||||
ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout)
|
ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(lookupEnv)
|
listenAddr, upstreamURL, err := config.ListenAddrAndUpstreamURL(lookupEnv)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("invalid setting: %w", err)
|
return fmt.Errorf("invalid setting: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The settings have checked that the address has a port.
|
// The settings have checked that the address has a port.
|
||||||
_, port, _ := net.SplitHostPort(cfg.ListenAddr)
|
_, port, _ := net.SplitHostPort(listenAddr)
|
||||||
health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath
|
health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody)
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody)
|
||||||
@@ -75,7 +77,7 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
|
|||||||
return fmt.Errorf("%w %s", errHealthEndpoint, res.Status)
|
return fmt.Errorf("%w %s", errHealthEndpoint, res.Status)
|
||||||
}
|
}
|
||||||
|
|
||||||
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(cfg.UpstreamURL))
|
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(upstreamURL))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("connect to the app: %w", err)
|
return fmt.Errorf("connect to the app: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@@ -24,12 +26,15 @@ func TestHealthCheck(t *testing.T) {
|
|||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
exited := make(chan int, 1)
|
exited := make(chan int, 1)
|
||||||
|
settings := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: app.URL,
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
exited <- run(ctx, map[string]string{
|
exited <- run(ctx, settings, out)
|
||||||
listenAddr: localhost + ":0",
|
|
||||||
upstreamURL: app.URL,
|
|
||||||
}, out)
|
|
||||||
}()
|
}()
|
||||||
|
|
||||||
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
||||||
@@ -40,6 +45,21 @@ func TestHealthCheck(t *testing.T) {
|
|||||||
|
|
||||||
wantHealthCheck(t, env, 0, "")
|
wantHealthCheck(t, env, 0, "")
|
||||||
|
|
||||||
|
// The health check reads those two settings alone, here given as
|
||||||
|
// files: a removed or invalid token file, or an invalid value of
|
||||||
|
// another setting, does not fail it.
|
||||||
|
for _, other := range []struct{ name, value string }{
|
||||||
|
{"SWWAF_METRICS_TOKEN_FILE", filepath.Join(t.TempDir(), "removed")},
|
||||||
|
{"SWWAF_METRICS_TOKEN_FILE", writeFile(t, "too short\n")},
|
||||||
|
{"SWWAF_MODE", "neither"},
|
||||||
|
} {
|
||||||
|
wantHealthCheck(t, map[string]string{
|
||||||
|
listenAddr + "_FILE": writeFile(t, ":"+port+"\n"),
|
||||||
|
upstreamURL + "_FILE": writeFile(t, app.URL+"\n"),
|
||||||
|
other.name: other.value,
|
||||||
|
}, 0, "")
|
||||||
|
}
|
||||||
|
|
||||||
app.Close()
|
app.Close()
|
||||||
wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ")
|
wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ")
|
||||||
|
|
||||||
@@ -95,3 +115,18 @@ func wantHealthCheck(t *testing.T, env map[string]string, status int, message st
|
|||||||
got, wrote, status, message)
|
got, wrote, status, message)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// writeFile writes contents to a file in a directory of its own, removed
|
||||||
|
// when the test ends, and returns the file's path.
|
||||||
|
func writeFile(t *testing.T, contents string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "setting")
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte(contents), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
||||||
// serves requests until it is told to stop, and then stops in an orderly
|
// the rule files and the state files, serves requests until it is told to
|
||||||
// way.
|
// stop, and then stops in an orderly way, writing the state files.
|
||||||
package smallwebwaf
|
package smallwebwaf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -18,7 +18,10 @@ import (
|
|||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/state"
|
||||||
)
|
)
|
||||||
|
|
||||||
// shutdownTimeout is how long requests in progress may take to finish
|
// shutdownTimeout is how long requests in progress may take to finish
|
||||||
@@ -26,6 +29,11 @@ import (
|
|||||||
// runit and docker wait a little longer before they kill the process.
|
// runit and docker wait a little longer before they kill the process.
|
||||||
const shutdownTimeout = 5 * time.Second
|
const shutdownTimeout = 5 * time.Second
|
||||||
|
|
||||||
|
// remoteLogStopTimeout is how long, as smallwebwaf stops, the log lines
|
||||||
|
// still waiting are sent to SWWAF_LOG_REMOTE_URL before they are given
|
||||||
|
// up. stdout has carried them.
|
||||||
|
const remoteLogStopTimeout = 2 * time.Second
|
||||||
|
|
||||||
// Params are what Run needs from the process.
|
// Params are what Run needs from the process.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
// Version is the version of the binary, set when it is built.
|
// Version is the version of the binary, set when it is built.
|
||||||
@@ -55,8 +63,9 @@ func Main(version string) int {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run reads the settings, then serves requests until ctx is done. It
|
// Run reads the settings, the rule files and the state files, then serves
|
||||||
// returns the process's exit status, 1 when smallwebwaf cannot start.
|
// 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 {
|
func Run(ctx context.Context, params Params) int {
|
||||||
processLog := requestlog.NewProcessLogger(params.Stdout)
|
processLog := requestlog.NewProcessLogger(params.Stdout)
|
||||||
|
|
||||||
@@ -67,6 +76,64 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// While SWWAF_LOG_REMOTE_URL is set, every line on stdout from here on
|
||||||
|
// is sent there too.
|
||||||
|
stdout := params.Stdout
|
||||||
|
|
||||||
|
var remote *remotelog.Sender
|
||||||
|
|
||||||
|
if cfg.LogRemoteURL != nil {
|
||||||
|
remote = newRemoteLogSender(cfg)
|
||||||
|
stdout = io.MultiWriter(params.Stdout, remote)
|
||||||
|
processLog = requestlog.NewProcessLogger(stdout)
|
||||||
|
|
||||||
|
stopSending := startSending(ctx, remote, processLog)
|
||||||
|
defer stopSending()
|
||||||
|
}
|
||||||
|
|
||||||
|
ruleFiles, err := rules.Load(rules.Params{
|
||||||
|
Dir: cfg.RulesDir,
|
||||||
|
Enabled: cfg.RulesEnabled,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
processLog.Error("cannot use the rule files", "error", err.Error())
|
||||||
|
|
||||||
|
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: stdout,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
GeoJSURL: lookup.URL,
|
||||||
|
Now: now,
|
||||||
|
Rules: ruleFiles,
|
||||||
|
})
|
||||||
|
if remote != nil {
|
||||||
|
server.Metrics.AddRemoteLog(remote)
|
||||||
|
}
|
||||||
|
|
||||||
|
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,
|
||||||
|
Metrics: server.Metrics,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
processLog.Error("cannot use the state files", "error", err.Error())
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
||||||
@@ -75,26 +142,58 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: params.Stdout,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: lookup.URL,
|
|
||||||
})
|
|
||||||
|
|
||||||
processLog.Info("starting",
|
processLog.Info("starting",
|
||||||
"version", params.Version,
|
"version", params.Version,
|
||||||
"address", listener.Addr().String(),
|
"address", listener.Addr().String(),
|
||||||
"settings", cfg)
|
"settings", cfg)
|
||||||
|
|
||||||
return serve(ctx, server, listener, processLog)
|
return serve(ctx, server.Server, listener, files, ruleFiles, processLog)
|
||||||
}
|
}
|
||||||
|
|
||||||
// serve serves requests on listener until ctx is done, then gives the
|
// newRemoteLogSender returns a sender of the log lines to
|
||||||
// requests in progress shutdownTimeout to finish.
|
// SWWAF_LOG_REMOTE_URL, with the settings for it.
|
||||||
|
func newRemoteLogSender(cfg *config.Config) *remotelog.Sender {
|
||||||
|
return remotelog.New(remotelog.Params{
|
||||||
|
URL: cfg.LogRemoteURL,
|
||||||
|
RootCAs: cfg.LogRemoteTLSCAs,
|
||||||
|
Buffer: cfg.LogRemoteBuffer,
|
||||||
|
Facility: cfg.LogRemoteFacility,
|
||||||
|
AppName: cfg.LogRemoteAppName,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// startSending runs remote until the function it returns is called, which
|
||||||
|
// then waits at most remoteLogStopTimeout for the lines still waiting to
|
||||||
|
// be sent. Sending goes on after ctx is done, so that the lines written
|
||||||
|
// while smallwebwaf stops are sent too.
|
||||||
|
func startSending(
|
||||||
|
ctx context.Context, remote *remotelog.Sender, processLog *slog.Logger,
|
||||||
|
) func() {
|
||||||
|
sending, stop := context.WithCancel(context.WithoutCancel(ctx))
|
||||||
|
sent := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
remote.Run(sending, processLog)
|
||||||
|
close(sent)
|
||||||
|
}()
|
||||||
|
|
||||||
|
return func() {
|
||||||
|
stop()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-sent:
|
||||||
|
case <-time.After(remoteLogStopTimeout):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// serve serves requests on listener, writes the state files as they are
|
||||||
|
// due, takes in an admin's edits of them, and reads the rule files again
|
||||||
|
// as they change, until ctx is done. Then it gives the requests in
|
||||||
|
// progress shutdownTimeout to finish, and writes every state file.
|
||||||
func serve(
|
func serve(
|
||||||
ctx context.Context, server *http.Server, listener net.Listener,
|
ctx context.Context, server *http.Server, listener net.Listener,
|
||||||
processLog *slog.Logger,
|
files *state.Files, ruleFiles *rules.Files, processLog *slog.Logger,
|
||||||
) int {
|
) int {
|
||||||
served := make(chan error, 1)
|
served := make(chan error, 1)
|
||||||
|
|
||||||
@@ -102,6 +201,28 @@ func serve(
|
|||||||
served <- server.Serve(listener)
|
served <- server.Serve(listener)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
writing, stopWriting := context.WithCancel(ctx)
|
||||||
|
defer stopWriting()
|
||||||
|
|
||||||
|
written := make(chan struct{})
|
||||||
|
watched := make(chan struct{})
|
||||||
|
rulesWatched := make(chan struct{})
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.Run(writing)
|
||||||
|
close(written)
|
||||||
|
}()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
files.Watch(writing)
|
||||||
|
close(watched)
|
||||||
|
}()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
ruleFiles.Watch(writing)
|
||||||
|
close(rulesWatched)
|
||||||
|
}()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-served:
|
case err := <-served:
|
||||||
processLog.Error("serving failed", "error", err.Error())
|
processLog.Error("serving failed", "error", err.Error())
|
||||||
@@ -131,6 +252,24 @@ func serve(
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Run and Watch have ended, so nothing else reads or 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
|
||||||
|
<-watched
|
||||||
|
<-rulesWatched
|
||||||
|
|
||||||
|
err = files.WriteAll()
|
||||||
|
if err != nil {
|
||||||
|
processLog.Error("writing the state files failed", "error", err.Error())
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
processLog.Info("stopped")
|
processLog.Info("stopped")
|
||||||
|
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
package smallwebwaf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log/slog"
|
||||||
|
"net/url"
|
||||||
|
"testing"
|
||||||
|
"testing/synctest"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The stop's tests run in a synctest bubble, where the time package runs
|
||||||
|
// on a clock of the test's own, so that how long the stop takes can be
|
||||||
|
// told exactly. The sender is held up by its process log, not by the
|
||||||
|
// network: a goroutine of the bubble that waits on the network keeps that
|
||||||
|
// clock from moving on.
|
||||||
|
|
||||||
|
func TestStopWaitsForTheSenderToFinish(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
took := stopHeldSender(t, time.Second)
|
||||||
|
if took != time.Second {
|
||||||
|
t.Errorf("the stop took %s, want the second the sender took", took)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStopWaitsForTheSenderAtMostTwoSeconds(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
synctest.Test(t, func(t *testing.T) {
|
||||||
|
took := stopHeldSender(t, time.Minute)
|
||||||
|
if took != 2*time.Second {
|
||||||
|
t.Errorf("the stop took %s, want 2s", took)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// heldLog holds each line written to it until it is closed.
|
||||||
|
type heldLog chan struct{}
|
||||||
|
|
||||||
|
// Write waits until the log is closed.
|
||||||
|
func (l heldLog) Write(p []byte) (int, error) {
|
||||||
|
<-l
|
||||||
|
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopHeldSender starts sending to an endpoint the sender cannot connect
|
||||||
|
// to, holds the sender as it logs that failure until release has passed,
|
||||||
|
// stops the sending, and returns how long the stop took. It returns once
|
||||||
|
// the sender has ended, as a bubble must.
|
||||||
|
func stopHeldSender(t *testing.T, release time.Duration) time.Duration {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
log := make(heldLog)
|
||||||
|
sender := remotelog.New(remotelog.Params{
|
||||||
|
// No port is 65536, so each attempt to connect fails at once,
|
||||||
|
// before it reaches the network.
|
||||||
|
URL: &url.URL{Scheme: remotelog.SchemeTCP, Host: "127.0.0.1:65536"},
|
||||||
|
Buffer: 1,
|
||||||
|
})
|
||||||
|
|
||||||
|
stopSending := startSending(t.Context(), sender,
|
||||||
|
slog.New(slog.NewJSONHandler(log, nil)))
|
||||||
|
|
||||||
|
synctest.Wait()
|
||||||
|
time.AfterFunc(release, func() { close(log) })
|
||||||
|
|
||||||
|
stopped := time.Now()
|
||||||
|
|
||||||
|
stopSending()
|
||||||
|
|
||||||
|
took := time.Since(stopped)
|
||||||
|
|
||||||
|
time.Sleep(release)
|
||||||
|
synctest.Wait()
|
||||||
|
|
||||||
|
return took
|
||||||
|
}
|
||||||
@@ -8,6 +8,10 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -24,9 +28,17 @@ const (
|
|||||||
// testVersion is the version the tests give smallwebwaf.
|
// testVersion is the version the tests give smallwebwaf.
|
||||||
testVersion = "test"
|
testVersion = "test"
|
||||||
// localhost is where the tests listen.
|
// localhost is where the tests listen.
|
||||||
localhost = "127.0.0.1"
|
localhost = "127.0.0.1"
|
||||||
listenAddr = "SWWAF_LISTEN_ADDR"
|
listenAddr = "SWWAF_LISTEN_ADDR"
|
||||||
upstreamURL = "SWWAF_UPSTREAM_URL"
|
upstreamURL = "SWWAF_UPSTREAM_URL"
|
||||||
|
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||||
|
stateDir = "SWWAF_STATE_DIR"
|
||||||
|
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
|
||||||
|
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
||||||
|
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||||
|
rulesDir = "SWWAF_RULES_DIR"
|
||||||
|
// greeting is what the tests' app answers.
|
||||||
|
greeting = "hello from the app"
|
||||||
)
|
)
|
||||||
|
|
||||||
// output collects what smallwebwaf writes on stdout.
|
// output collects what smallwebwaf writes on stdout.
|
||||||
@@ -69,11 +81,19 @@ func (o *output) line(t *testing.T, key, value string) map[string]any {
|
|||||||
time.Sleep(pollInterval)
|
time.Sleep(pollInterval)
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.buf.String())
|
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.text())
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// text returns everything written so far.
|
||||||
|
func (o *output) text() string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
return o.buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
// run runs smallwebwaf with the settings in env until ctx is done, and
|
// run runs smallwebwaf with the settings in env until ctx is done, and
|
||||||
// returns its exit status.
|
// returns its exit status.
|
||||||
func run(ctx context.Context, env map[string]string, out *output) int {
|
func run(ctx context.Context, env map[string]string, out *output) int {
|
||||||
@@ -107,6 +127,28 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
|
||||||
|
|
||||||
|
out := &output{}
|
||||||
|
|
||||||
|
status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out)
|
||||||
|
if status != 1 {
|
||||||
|
t.Errorf("exit status %d, want 1", status)
|
||||||
|
}
|
||||||
|
|
||||||
|
line := out.line(t, "msg", "invalid setting")
|
||||||
|
if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" {
|
||||||
|
t.Errorf("start refused with %v", line)
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(out.text(), token) {
|
||||||
|
t.Errorf("the output shows the token:\n%s", out.text())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestAddressInUseStopsTheStart(t *testing.T) {
|
func TestAddressInUseStopsTheStart(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -121,7 +163,11 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
|
|||||||
|
|
||||||
out := &output{}
|
out := &output{}
|
||||||
|
|
||||||
status := run(t.Context(), map[string]string{listenAddr: taken.Addr().String()}, out)
|
status := run(t.Context(), map[string]string{
|
||||||
|
listenAddr: taken.Addr().String(),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
}, out)
|
||||||
if status != 1 {
|
if status != 1 {
|
||||||
t.Errorf("exit status %d, want 1", status)
|
t.Errorf("exit status %d, want 1", status)
|
||||||
}
|
}
|
||||||
@@ -132,11 +178,8 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
|
|||||||
func TestServesUntilToldToStop(t *testing.T) {
|
func TestServesUntilToldToStop(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
app := httptest.NewServer(http.HandlerFunc(
|
appURL := startApp(t)
|
||||||
func(w http.ResponseWriter, _ *http.Request) {
|
dir := t.TempDir()
|
||||||
_, _ = io.WriteString(w, "hello from the app")
|
|
||||||
}))
|
|
||||||
defer app.Close()
|
|
||||||
|
|
||||||
ctx, stop := context.WithCancel(t.Context())
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
out := &output{}
|
out := &output{}
|
||||||
@@ -145,12 +188,19 @@ func TestServesUntilToldToStop(t *testing.T) {
|
|||||||
go func() {
|
go func() {
|
||||||
exited <- run(ctx, map[string]string{
|
exited <- run(ctx, map[string]string{
|
||||||
listenAddr: localhost + ":0",
|
listenAddr: localhost + ":0",
|
||||||
upstreamURL: app.URL,
|
upstreamURL: appURL,
|
||||||
|
stateDir: dir,
|
||||||
|
rulesDir: filepath.Join("..", "..", "share", "rules.d"),
|
||||||
}, out)
|
}, out)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
// The default rule file is read.
|
||||||
|
if rules := out.line(t, "msg", "read the rule files")["rules"]; rules != 12.0 {
|
||||||
|
t.Errorf("read %v rules from the default rule file, want 12", rules)
|
||||||
|
}
|
||||||
|
|
||||||
starting := out.line(t, "msg", "starting")
|
starting := out.line(t, "msg", "starting")
|
||||||
wantStartingLine(t, starting, app.URL)
|
wantStartingLine(t, starting, appURL, dir)
|
||||||
|
|
||||||
addr, _ := starting["address"].(string)
|
addr, _ := starting["address"].(string)
|
||||||
wantGreeting(t, "http://"+addr+"/")
|
wantGreeting(t, "http://"+addr+"/")
|
||||||
@@ -170,30 +220,436 @@ func TestServesUntilToldToStop(t *testing.T) {
|
|||||||
out.line(t, "msg", "stopped")
|
out.line(t, "msg", "stopped")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStateKeptAcrossRestarts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
rateLimitPerDay: "2",
|
||||||
|
// Neither comes due in the test: the files are written as
|
||||||
|
// smallwebwaf stops.
|
||||||
|
stateWriteDelay: "1h",
|
||||||
|
stateCounterInterval: "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(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
trustedProxies: 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 TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const (
|
||||||
|
// bans.json as an admin writes it with a ban, permanent, on
|
||||||
|
// 203.0.113.0/24, and with none.
|
||||||
|
oneBan = `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", ` +
|
||||||
|
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`
|
||||||
|
noBan = `{"version": 1, "bans": []}`
|
||||||
|
)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: dir,
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
trustedProxies: localhost + "/32",
|
||||||
|
// No write comes due in the test, so only the watch on the
|
||||||
|
// directory can take the edits in.
|
||||||
|
stateWriteDelay: "1h",
|
||||||
|
stateCounterInterval: "1h",
|
||||||
|
}
|
||||||
|
|
||||||
|
runUntilStopped(t, env, func(url string) {
|
||||||
|
path := filepath.Join(dir, "bans.json")
|
||||||
|
|
||||||
|
saveUntilAnswered(t, path, oneBan, url, "203.0.113.9", http.StatusForbidden)
|
||||||
|
wantStatus(t, url, "198.51.100.7", http.StatusOK)
|
||||||
|
saveUntilAnswered(t, path, noBan, url, "203.0.113.9", http.StatusOK)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRuleFileAddedWhileRunningTakesEffect(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: dir,
|
||||||
|
// The requests sent until the rule takes effect must not break a
|
||||||
|
// rate limit, whose ban would refuse them too.
|
||||||
|
"SWWAF_RATE_LIMIT_PER_MINUTE": "off",
|
||||||
|
}
|
||||||
|
|
||||||
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
|
wantGreeting(t, url)
|
||||||
|
|
||||||
|
// Written once: each change would start the rule files' wait
|
||||||
|
// again. A file written before smallwebwaf watches the directory is
|
||||||
|
// read once it does.
|
||||||
|
err := os.WriteFile(filepath.Join(dir, "50-app.rules"),
|
||||||
|
[]byte("everything path block ^/\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write the rule file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// As long as that takes, so that a slow test process cannot fail
|
||||||
|
// the test.
|
||||||
|
for statusFrom(t, url, "203.0.113.9") != http.StatusForbidden {
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
out.line(t, "action", "rule_blocked")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRuleFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "00-default.rules")
|
||||||
|
|
||||||
|
err := os.WriteFile(path, []byte("# probes\nenv-file path bann ^/\\.env$\n"), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write the rule file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
wantRulesRefused(t, dir, path+`, line 2: the action "bann" is not log, block or ban`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := filepath.Join(t.TempDir(), "rules.d")
|
||||||
|
|
||||||
|
wantRulesRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
|
||||||
|
": no such file or directory")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
endpoint, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = endpoint.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
"SWWAF_LOG_REMOTE_URL": "syslog+tcp://" + endpoint.Addr().String(),
|
||||||
|
}
|
||||||
|
|
||||||
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
|
wantGreeting(t, url)
|
||||||
|
})
|
||||||
|
out.line(t, "type", "request")
|
||||||
|
|
||||||
|
// smallwebwaf connected as it started, and closes the connection once
|
||||||
|
// it has sent the lines written as it stopped.
|
||||||
|
conn, err := endpoint.Accept()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("accept: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
received, err := io.ReadAll(conn)
|
||||||
|
_ = conn.Close()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Lines written at once by several goroutines may reach stdout and
|
||||||
|
// the endpoint in different orders.
|
||||||
|
sent := messages(t, string(received))
|
||||||
|
written := slices.Collect(strings.Lines(out.text()))
|
||||||
|
|
||||||
|
slices.Sort(sent)
|
||||||
|
slices.Sort(written)
|
||||||
|
|
||||||
|
if !slices.Equal(sent, written) {
|
||||||
|
t.Errorf("sent\n%v\nwrote\n%v", sent, written)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
const token = "0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
// The endpoint takes connections and never answers, so the TLS
|
||||||
|
// handshake of each waits on it, and no line is ever sent.
|
||||||
|
endpoint, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = endpoint.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
env := map[string]string{
|
||||||
|
listenAddr: localhost + ":0",
|
||||||
|
upstreamURL: startApp(t),
|
||||||
|
stateDir: t.TempDir(),
|
||||||
|
rulesDir: t.TempDir(),
|
||||||
|
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
|
||||||
|
"SWWAF_LOG_REMOTE_BUFFER": "1",
|
||||||
|
"SWWAF_METRICS_TOKEN": token,
|
||||||
|
}
|
||||||
|
|
||||||
|
out := runUntilStopped(t, env, func(url string) {
|
||||||
|
wantGreeting(t, url)
|
||||||
|
|
||||||
|
// More than one line has been written, and the buffer holds the
|
||||||
|
// last.
|
||||||
|
metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
|
||||||
|
for _, series := range []string{
|
||||||
|
"smallwebwaf_remote_log_lines_sent_total 0",
|
||||||
|
"smallwebwaf_remote_log_buffer_depth 1",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(metrics, "\n"+series+"\n") {
|
||||||
|
t.Errorf("no %q in the metrics:\n%s", series, metrics)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") ||
|
||||||
|
!strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") {
|
||||||
|
t.Errorf("no line dropped in the metrics:\n%s", metrics)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Closed, the endpoint refuses the connection made to send the
|
||||||
|
// lines still waiting at the stop, which then does not wait.
|
||||||
|
_ = endpoint.Close()
|
||||||
|
})
|
||||||
|
|
||||||
|
out.line(t, "type", "request")
|
||||||
|
}
|
||||||
|
|
||||||
|
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, rulesDir: t.TempDir(),
|
||||||
|
}, 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// wantRulesRefused runs smallwebwaf with its rule files in dir, and
|
||||||
|
// checks that it stops at start, with the error want. If it starts
|
||||||
|
// instead, it is stopped after waitLimit.
|
||||||
|
func wantRulesRefused(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: t.TempDir(), rulesDir: dir,
|
||||||
|
}, out)
|
||||||
|
if status != 1 {
|
||||||
|
t.Fatalf("exit status %d, want 1", status)
|
||||||
|
}
|
||||||
|
|
||||||
|
line := out.line(t, "msg", "cannot use the rule files")
|
||||||
|
if line["error"] != want || line["level"] != "ERROR" {
|
||||||
|
t.Errorf("start refused with %v, want the error %q", line, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// startApp starts an app that answers every request with greeting, and
|
||||||
|
// returns its URL.
|
||||||
|
func startApp(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
app := httptest.NewServer(http.HandlerFunc(
|
||||||
|
func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = io.WriteString(w, greeting)
|
||||||
|
}))
|
||||||
|
t.Cleanup(app.Close)
|
||||||
|
|
||||||
|
return app.URL
|
||||||
|
}
|
||||||
|
|
||||||
|
// runUntilStopped runs smallwebwaf with the settings in env, has use send
|
||||||
|
// it requests at url, then stops it as SIGTERM does, checks that it
|
||||||
|
// stopped in order, and returns its output.
|
||||||
|
func runUntilStopped(
|
||||||
|
t *testing.T, env map[string]string, use func(url string),
|
||||||
|
) *output {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
ctx, stop := context.WithCancel(t.Context())
|
||||||
|
out := &output{}
|
||||||
|
exited := make(chan int, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
exited <- run(ctx, env, out)
|
||||||
|
}()
|
||||||
|
|
||||||
|
addr, _ := out.line(t, "msg", "starting")["address"].(string)
|
||||||
|
use("http://" + addr + "/")
|
||||||
|
stop()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case status := <-exited:
|
||||||
|
if status != 0 {
|
||||||
|
t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text())
|
||||||
|
}
|
||||||
|
case <-time.After(waitLimit):
|
||||||
|
t.Fatal("still running after being told to stop")
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
// wantStartingLine checks that the line at start gives the version and
|
// wantStartingLine checks that the line at start gives the version and
|
||||||
// every setting's value.
|
// every setting's value.
|
||||||
func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
|
func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
settings, _ := line["settings"].(map[string]any)
|
settings, _ := line["settings"].(map[string]any)
|
||||||
want := map[string]any{
|
want := map[string]any{
|
||||||
listenAddr: localhost + ":0",
|
listenAddr: localhost + ":0",
|
||||||
upstreamURL: appURL,
|
upstreamURL: appURL,
|
||||||
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
stateDir: dir,
|
||||||
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
|
"SWWAF_MODE": "enforce",
|
||||||
"SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m",
|
stateWriteDelay: "10s",
|
||||||
"SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s",
|
stateCounterInterval: "15m",
|
||||||
"SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m",
|
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
||||||
"SWWAF_REQUEST_MAX_BYTES": "100M",
|
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
|
||||||
"SWWAF_RESPONSE_MAX_BYTES": "5G",
|
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
|
||||||
"SWWAF_ALLOW_NETS": "",
|
"SWWAF_CLIENT_IDLE_TIMEOUT": "120s",
|
||||||
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
|
"SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m",
|
||||||
"SWWAF_DENY_NETS": "",
|
"SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s",
|
||||||
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
|
"SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m",
|
||||||
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
|
"SWWAF_REQUEST_MAX_BYTES": "100M",
|
||||||
"SWWAF_RATE_LIMIT_PER_DAY": "50000",
|
"SWWAF_RESPONSE_MAX_BYTES": "5G",
|
||||||
"SWWAF_DENIED_COUNTRIES": "",
|
"SWWAF_ALLOW_NETS": "",
|
||||||
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
|
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
|
||||||
|
"SWWAF_DENY_NETS": "",
|
||||||
|
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
|
||||||
|
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
|
||||||
|
rateLimitPerDay: "50000",
|
||||||
|
"SWWAF_RATE_LIMIT_EXEMPT_PATHS": "",
|
||||||
|
"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_ATTACK_BAN_DURATION": "7d",
|
||||||
|
"SWWAF_MAX_BANS": "5000",
|
||||||
|
"SWWAF_BAN_SCOPE_V4_PREFIX": "32",
|
||||||
|
"SWWAF_RULES_ENABLED": "true",
|
||||||
}
|
}
|
||||||
|
|
||||||
for name, value := range want {
|
for name, value := range want {
|
||||||
@@ -228,7 +684,156 @@ func wantGreeting(t *testing.T, url string) {
|
|||||||
body, err := io.ReadAll(res.Body)
|
body, err := io.ReadAll(res.Body)
|
||||||
_ = res.Body.Close()
|
_ = res.Body.Close()
|
||||||
|
|
||||||
if err != nil || string(body) != "hello from the app" {
|
if err != nil || string(body) != greeting {
|
||||||
t.Errorf("got %q (%v), want the app's answer", body, err)
|
t.Errorf("got %q (%v), want the app's answer", body, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// messages returns the message of each record in received, octet-counted
|
||||||
|
// frames of RFC 5424 records with the default facility and app name, each
|
||||||
|
// with the newline that ends a line on stdout.
|
||||||
|
func messages(t *testing.T, received string) []string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
hostname, _ := os.Hostname()
|
||||||
|
header := " " + hostname + " " + hostname + " - - - "
|
||||||
|
|
||||||
|
var found []string
|
||||||
|
|
||||||
|
for received != "" {
|
||||||
|
count, rest, _ := strings.Cut(received, " ")
|
||||||
|
|
||||||
|
length, err := strconv.Atoi(count)
|
||||||
|
if err != nil || length > len(rest) {
|
||||||
|
t.Fatalf("no frame at %q", received)
|
||||||
|
}
|
||||||
|
|
||||||
|
record := rest[:length]
|
||||||
|
received = rest[length:]
|
||||||
|
|
||||||
|
_, message, ok := strings.Cut(record, header)
|
||||||
|
if !ok || !strings.HasPrefix(record, "<134>1 ") {
|
||||||
|
t.Fatalf("record %q, want priority <134> and header %q", record, header)
|
||||||
|
}
|
||||||
|
|
||||||
|
found = append(found, message+"\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
return found
|
||||||
|
}
|
||||||
|
|
||||||
|
// metricsText asks for the metrics at url with token, and returns them.
|
||||||
|
func metricsText(t *testing.T, url, token string) string {
|
||||||
|
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("Authorization", "Bearer "+token)
|
||||||
|
|
||||||
|
transport := &http.Transport{}
|
||||||
|
defer transport.CloseIdleConnections()
|
||||||
|
|
||||||
|
res, err := (&http.Client{Transport: transport}).Do(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("request: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := io.ReadAll(res.Body)
|
||||||
|
_ = res.Body.Close()
|
||||||
|
|
||||||
|
if err != nil || res.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("metrics answered %d (%v)", res.StatusCode, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return string(body)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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()
|
||||||
|
|
||||||
|
got := statusFrom(t, url, from)
|
||||||
|
if got != status {
|
||||||
|
t.Errorf("request from %s: status %d, want %d", from, got, status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// saveUntilAnswered writes content to the state file at path, as an
|
||||||
|
// admin saves an edit of it, until a request to url from the client at
|
||||||
|
// from is answered with status. The file is written again before each
|
||||||
|
// request, since smallwebwaf may not watch its directory yet when it is
|
||||||
|
// first written. It waits as long as that takes, so that a slow test
|
||||||
|
// process cannot fail the test.
|
||||||
|
func saveUntilAnswered(t *testing.T, path, content, url, from string, status int) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
for {
|
||||||
|
err := os.WriteFile(path, []byte(content), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write %s: %v", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if statusFrom(t, url, from) == status {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(pollInterval)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// statusFrom returns the status a request to url from the client at
|
||||||
|
// from, as X-Forwarded-For names it, is answered with.
|
||||||
|
func statusFrom(t *testing.T, url, from string) int {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
|
||||||
|
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()
|
||||||
|
|
||||||
|
return res.StatusCode
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,728 @@
|
|||||||
|
// 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,
|
||||||
|
// Watch takes in an admin's edit of one while smallwebwaf runs, and Run
|
||||||
|
// and WriteAll write them. The disk is read and written outside the
|
||||||
|
// parts' locks, which are held only to take a snapshot or to put in what
|
||||||
|
// a file holds, so that no request waits on the disk.
|
||||||
|
package state
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/fsnotify/fsnotify"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||||
|
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||||
|
"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")
|
||||||
|
errCause = errors.New("is not limit, attack or admin")
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 taken in, the edits set aside,
|
||||||
|
// and the writes that fail.
|
||||||
|
ProcessLog *slog.Logger
|
||||||
|
// Metrics count each file's writes, and the edits taken in and set
|
||||||
|
// aside.
|
||||||
|
Metrics *metrics.Metrics
|
||||||
|
}
|
||||||
|
|
||||||
|
// Files are the state files of a running smallwebwaf.
|
||||||
|
type Files struct {
|
||||||
|
params Params
|
||||||
|
|
||||||
|
// mu is held while a file is read for an edit, and while it is
|
||||||
|
// written, so that Watch and the writes take turns. No request takes
|
||||||
|
// it.
|
||||||
|
mu sync.Mutex
|
||||||
|
// sums are the SHA-256 sums of what each file held, by name, when
|
||||||
|
// smallwebwaf last read or wrote it. A file that holds anything else
|
||||||
|
// has been edited since.
|
||||||
|
sums map[string][sha256.Size]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// bansFile is bans.json, indented for an admin to read and edit.
|
||||||
|
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, a ban an admin added may have no cause, which makes it an
|
||||||
|
// admin's, and lifted is left out until an admin lifts the ban.
|
||||||
|
type banEntry struct {
|
||||||
|
Netblock netip.Prefix `json:"netblock"`
|
||||||
|
Start time.Time `json:"start"`
|
||||||
|
Expires *time.Time `json:"expires"`
|
||||||
|
Cause string `json:"cause"`
|
||||||
|
Reason string `json:"reason,omitempty"`
|
||||||
|
Lifted *time.Time `json:"lifted,omitempty"`
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
f := &Files{params: params, sums: map[string][sha256.Size]byte{}}
|
||||||
|
|
||||||
|
bansRead, bansErr := f.read(bansJSON)
|
||||||
|
clientsRead, clientsErr := f.read(clientsJSON)
|
||||||
|
lookupsRead, lookupsErr := f.read(lookupsJSON)
|
||||||
|
|
||||||
|
err = errors.Join(bansErr, clientsErr, lookupsErr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
||||||
|
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
|
||||||
|
|
||||||
|
return f, 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. Each write takes in an admin's edit of its file first,
|
||||||
|
// as writeFile describes.
|
||||||
|
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.writeFile(bansJSON))
|
||||||
|
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.writeFile(bansJSON), f.writeFile(clientsJSON),
|
||||||
|
f.writeFile(lookupsJSON))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
|
||||||
|
// state file as soon as it is saved: what the file holds replaces what
|
||||||
|
// smallwebwaf held for it. An edit that does not parse is left for the
|
||||||
|
// file's next write, which sets it aside, since a file can be read while
|
||||||
|
// an editor is still writing it. If Dir cannot be watched, that is
|
||||||
|
// logged, and an edit is taken in only before its file is written.
|
||||||
|
func (f *Files) Watch(ctx context.Context) {
|
||||||
|
watcher, err := fsnotify.NewWatcher()
|
||||||
|
if err == nil {
|
||||||
|
defer func() {
|
||||||
|
_ = watcher.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
err = watcher.Add(f.params.Dir)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error("cannot watch the state files for edits",
|
||||||
|
"error", err.Error())
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.ProcessLog.Info("watching the state files for edits",
|
||||||
|
"directory", f.params.Dir)
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case event := <-watcher.Events:
|
||||||
|
switch name := filepath.Base(event.Name); name {
|
||||||
|
case bansJSON, clientsJSON, lookupsJSON:
|
||||||
|
f.fileChanged(name)
|
||||||
|
}
|
||||||
|
case err = <-watcher.Errors:
|
||||||
|
f.params.ProcessLog.Warn("watching the state files failed",
|
||||||
|
"error", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// logFailure logs a write that failed.
|
||||||
|
func (f *Files) logFailure(err error) {
|
||||||
|
if err != nil {
|
||||||
|
f.params.ProcessLog.Error("writing the state files failed",
|
||||||
|
"error", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// fileChanged takes in what the state file name holds, as Watch sees it
|
||||||
|
// change, if that is an edit made since smallwebwaf last read or wrote
|
||||||
|
// the file. A file that cannot be read or does not parse is left for its
|
||||||
|
// next write.
|
||||||
|
func (f *Files) fileChanged(name string) {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
data, changed, err := f.readChanged(name)
|
||||||
|
if err != nil || !changed {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = f.takeInEdit(name, data)
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeInEdit takes in data, an edit of the state file name, as takeIn
|
||||||
|
// does, and counts and logs it. Every edit taken in while smallwebwaf
|
||||||
|
// runs, by Watch or by a write, is taken in here. An edit that does not
|
||||||
|
// parse is neither counted nor logged, and takeIn's error returned.
|
||||||
|
func (f *Files) takeInEdit(name string, data []byte) error {
|
||||||
|
_, err := f.takeIn(name, data, true)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Counted before it is logged, so that the count is there once the
|
||||||
|
// log line is.
|
||||||
|
f.params.Metrics.StateFileEditTakenIn(name)
|
||||||
|
f.params.ProcessLog.Info("took in an edit of a state file",
|
||||||
|
"file", filepath.Join(f.params.Dir, name))
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// read takes in the state file name at start, and returns how many
|
||||||
|
// entries it holds. A missing file holds none.
|
||||||
|
func (f *Files) read(name string) (int, error) {
|
||||||
|
data, changed, err := f.readChanged(name)
|
||||||
|
if err != nil || !changed {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return f.takeIn(name, data, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readChanged returns what the state file name holds, and whether that
|
||||||
|
// has changed since smallwebwaf last read or wrote the file, as it has
|
||||||
|
// for a file smallwebwaf never read or wrote. A missing file has not
|
||||||
|
// changed: it is written again at its next write.
|
||||||
|
func (f *Files) readChanged(name string) ([]byte, bool, error) {
|
||||||
|
path := filepath.Join(f.params.Dir, name)
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
|
||||||
|
if errors.Is(err, fs.ErrNotExist) {
|
||||||
|
return nil, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return data, sha256.Sum256(data) != f.sums[name], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeIn parses data, what the state file name holds, puts it into the
|
||||||
|
// part that keeps that state, in place of what the part held, and returns
|
||||||
|
// how many entries the file holds. edit is whether data is an admin's
|
||||||
|
// edit taken in while smallwebwaf runs, rather than the file read at the
|
||||||
|
// start. An error names the file and, where the JSON decoder tells it,
|
||||||
|
// the line and column, or else the entry.
|
||||||
|
func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
|
||||||
|
path := filepath.Join(f.params.Dir, name)
|
||||||
|
|
||||||
|
var entries int
|
||||||
|
|
||||||
|
switch name {
|
||||||
|
case bansJSON:
|
||||||
|
var file bansFile
|
||||||
|
|
||||||
|
err := parse(path, data, &file)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
held := make([]bans.Ban, 0, len(file.Bans))
|
||||||
|
for _, entry := range file.Bans {
|
||||||
|
held = append(held, entry.ban())
|
||||||
|
}
|
||||||
|
|
||||||
|
if edit {
|
||||||
|
f.params.Ledger.LoadEdit(held)
|
||||||
|
} else {
|
||||||
|
f.params.Ledger.Load(held)
|
||||||
|
}
|
||||||
|
|
||||||
|
entries = len(held)
|
||||||
|
case clientsJSON:
|
||||||
|
var file clientsFile
|
||||||
|
|
||||||
|
err := parse(path, data, &file)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.Limiter.Load(file.Clients, f.params.Now())
|
||||||
|
entries = len(file.Clients)
|
||||||
|
case lookupsJSON:
|
||||||
|
var file lookupsFile
|
||||||
|
|
||||||
|
err := parse(path, data, &file)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.GeoJS.Load(file.Lookups)
|
||||||
|
entries = len(file.Lookups)
|
||||||
|
}
|
||||||
|
|
||||||
|
f.sums[name] = sha256.Sum256(data)
|
||||||
|
|
||||||
|
return entries, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeFile writes the state file name from what smallwebwaf holds. An
|
||||||
|
// edit made since smallwebwaf last read or wrote the file is taken in
|
||||||
|
// first, so that it is not overwritten, or set aside if it does not
|
||||||
|
// parse. A file that cannot be read, or an edit that cannot be set
|
||||||
|
// aside, is left as it is, and the write given up. Every write is counted
|
||||||
|
// in the metrics, and one that fails or is given up as a failure.
|
||||||
|
func (f *Files) writeFile(name string) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
data, changed, err := f.readChanged(name)
|
||||||
|
if err == nil && changed {
|
||||||
|
err = f.takeInEdit(name, data)
|
||||||
|
if err != nil {
|
||||||
|
err = f.setAside(name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
data, err = f.encode(name)
|
||||||
|
if err != nil {
|
||||||
|
err = fmt.Errorf("encode %s: %w", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
err = write(f.params.Dir, name, data)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err == nil {
|
||||||
|
// The file holds data from here on, even if the directory sync
|
||||||
|
// fails, so that its next read does not take it for an admin's
|
||||||
|
// edit.
|
||||||
|
f.sums[name] = sha256.Sum256(data)
|
||||||
|
err = syncDirectory(f.params.Dir)
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.Metrics.StateFileWritten(name, len(data), err)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// setAside renames the state file name, an edit that does not parse with
|
||||||
|
// parseErr, to name.bad, for the admin to mend, and logs it with where in
|
||||||
|
// the file the error is. If the rename fails, the edit is left as it is,
|
||||||
|
// and the error returned is parseErr joined with the rename's.
|
||||||
|
func (f *Files) setAside(name string, parseErr error) error {
|
||||||
|
path := filepath.Join(f.params.Dir, name)
|
||||||
|
|
||||||
|
err := os.Rename(path, path+".bad")
|
||||||
|
if err != nil {
|
||||||
|
return errors.Join(parseErr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse",
|
||||||
|
"file", path+".bad", "error", parseErr.Error())
|
||||||
|
f.params.Metrics.StateFileEditSetAside(name)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// encode returns the state file name as smallwebwaf writes it, from a
|
||||||
|
// snapshot of the part that keeps that state.
|
||||||
|
func (f *Files) encode(name string) ([]byte, error) {
|
||||||
|
switch name {
|
||||||
|
case bansJSON:
|
||||||
|
held := f.params.Ledger.Snapshot()
|
||||||
|
|
||||||
|
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))}
|
||||||
|
for _, ban := range held {
|
||||||
|
file.Bans = append(file.Bans, newBanEntry(ban))
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.MarshalIndent(file, "", " ")
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return append(data, '\n'), nil
|
||||||
|
case clientsJSON:
|
||||||
|
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
||||||
|
default: // lookups.json
|
||||||
|
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newBanEntry returns ban as bans.json holds it.
|
||||||
|
func newBanEntry(ban bans.Ban) banEntry {
|
||||||
|
entry := banEntry{
|
||||||
|
Netblock: ban.Netblock, Start: ban.Start, Cause: ban.Cause, Reason: ban.Reason,
|
||||||
|
Notes: ban.Notes,
|
||||||
|
}
|
||||||
|
if !ban.Permanent() {
|
||||||
|
entry.Expires = &ban.Expires
|
||||||
|
}
|
||||||
|
|
||||||
|
if !ban.Lifted.IsZero() {
|
||||||
|
entry.Lifted = &ban.Lifted
|
||||||
|
}
|
||||||
|
|
||||||
|
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, Cause: e.Cause, Reason: e.Reason,
|
||||||
|
Notes: e.Notes,
|
||||||
|
}
|
||||||
|
if e.Expires != nil {
|
||||||
|
ban.Expires = *e.Expires
|
||||||
|
}
|
||||||
|
|
||||||
|
if e.Lifted != nil {
|
||||||
|
ban.Lifted = *e.Lifted
|
||||||
|
}
|
||||||
|
|
||||||
|
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. A cause other than limit,
|
||||||
|
// attack or admin, most likely misspelt, is refused too.
|
||||||
|
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")
|
||||||
|
case entry.Cause != "" && entry.Cause != bans.CauseLimit &&
|
||||||
|
entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin:
|
||||||
|
return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// parse reads data, what the state file at path holds, into file, a
|
||||||
|
// pointer to that file's struct, and checks its entries.
|
||||||
|
func parse(path string, data []byte, file stateFile) error {
|
||||||
|
// 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. syncDirectory must follow, 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// syncDirectory syncs dir to the disk, so that a rename in it lasts.
|
||||||
|
func syncDirectory(dir string) error {
|
||||||
|
directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
|
||||||
|
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())
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,35 @@
|
|||||||
|
package state
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The test is on write itself: a state file is read before it is
|
||||||
|
// written, and a directory in its place fails that read first.
|
||||||
|
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
// A directory named bans.json cannot be renamed over.
|
||||||
|
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("mkdir: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = write(dir, bansJSON, []byte("{}\n"))
|
||||||
|
if err == nil {
|
||||||
|
t.Error("writing over a directory did not fail")
|
||||||
|
}
|
||||||
|
|
||||||
|
entries, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read %s: %v", dir, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(entries) != 1 || entries[0].Name() != bansJSON {
|
||||||
|
t.Errorf("%s holds %v, want only bans.json", dir, entries)
|
||||||
|
}
|
||||||
|
}
|
||||||
+68
-13
@@ -1,12 +1,15 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# script/example-app: build the image, and on it the example app in
|
# script/example-app: build the image, and on it the example app in
|
||||||
# deploy/example-app, then run the app's container and check that the
|
# deploy/example-app, then run the app's container with a volume for the
|
||||||
# health check passes, that a request is served through smallwebwaf,
|
# state files and check that the health check passes, that a request is
|
||||||
# that `sv stop` stops smallwebwaf in order, and that `docker stop`
|
# served through smallwebwaf, that a second one in a minute bans the
|
||||||
# stops the container without having to kill it. The container and both
|
# client, that a probe for /.env bans another client, which its next
|
||||||
# images are removed however the script ends. Building the app needs
|
# request bans for good, that `sv stop` stops smallwebwaf in order, that
|
||||||
# network access, for nixpkgs' binary cache. script/check does not run
|
# `docker stop` stops the container without having to kill it, and that
|
||||||
# this.
|
# 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
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -18,9 +21,11 @@ NAME="$("$SCRIPT_DIR/projectname")-example-$$"
|
|||||||
IMAGE="$NAME-base"
|
IMAGE="$NAME-base"
|
||||||
APP_IMAGE="$NAME-app"
|
APP_IMAGE="$NAME-app"
|
||||||
CONTAINER="$NAME"
|
CONTAINER="$NAME"
|
||||||
|
VOLUME="$NAME-state"
|
||||||
|
|
||||||
cleanup() {
|
cleanup() {
|
||||||
docker rm --force "$CONTAINER" >/dev/null 2>&1 || true
|
docker rm --force "$CONTAINER" >/dev/null 2>&1 || true
|
||||||
|
docker volume rm --force "$VOLUME" >/dev/null 2>&1 || true
|
||||||
docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true
|
docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -48,9 +53,42 @@ healthy() {
|
|||||||
[ "$status" = healthy ]
|
[ "$status" = healthy ]
|
||||||
}
|
}
|
||||||
|
|
||||||
# logged <text>: the container's output holds text.
|
# logged <text>...: a line of the container's output holds every text,
|
||||||
|
# in any order.
|
||||||
logged() {
|
logged() {
|
||||||
docker logs "$CONTAINER" 2>&1 | grep -qF "$1"
|
lines="$(docker logs "$CONTAINER" 2>&1)"
|
||||||
|
for text in "$@"; do
|
||||||
|
lines="$(printf '%s\n' "$lines" | grep -F "$text")" || return 1
|
||||||
|
done
|
||||||
|
}
|
||||||
|
|
||||||
|
# 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 ]
|
||||||
|
}
|
||||||
|
|
||||||
|
# refused_from <client> <path>: a request for path from client, as
|
||||||
|
# X-Forwarded-For names it, gets 403. smallwebwaf believes the header
|
||||||
|
# from docker's gateway, a private address.
|
||||||
|
refused_from() {
|
||||||
|
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
|
||||||
|
--max-time 10 --header "X-Forwarded-For: $1" "http://$address$2")" || true
|
||||||
|
[ "$code" = 403 ]
|
||||||
}
|
}
|
||||||
|
|
||||||
main() {
|
main() {
|
||||||
@@ -62,18 +100,28 @@ main() {
|
|||||||
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
|
docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \
|
||||||
-t "$APP_IMAGE" deploy/example-app
|
-t "$APP_IMAGE" deploy/example-app
|
||||||
|
|
||||||
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
docker volume create "$VOLUME" >/dev/null
|
||||||
"$APP_IMAGE" >/dev/null
|
start_container
|
||||||
wait_for "the health check did not pass" healthy
|
|
||||||
echo "example-app: the health check passes"
|
echo "example-app: the health check passes"
|
||||||
|
|
||||||
address="$(docker port "$CONTAINER" 8080/tcp)"
|
|
||||||
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
|
page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" ||
|
||||||
fail "no answer on port 8080"
|
fail "no answer on port 8080"
|
||||||
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
|
[ "$page" = "hello from the example app" ] || fail "port 8080 answered $page"
|
||||||
wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"'
|
wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"'
|
||||||
echo "example-app: smallwebwaf passes a request to the app and its answer back"
|
echo "example-app: smallwebwaf passes a request to the app and its answer back"
|
||||||
|
|
||||||
|
refused || fail "a second request in a minute was not refused"
|
||||||
|
wait_for "smallwebwaf logged no ban" logged '"action":"rate_limited"'
|
||||||
|
echo "example-app: a second request in a minute bans the client"
|
||||||
|
|
||||||
|
refused_from 203.0.113.9 /.env || fail "a probe for /.env was not refused"
|
||||||
|
wait_for "smallwebwaf logged no ban for the probe" \
|
||||||
|
logged '"action":"banned"' '"rule_ids":["env-file"]'
|
||||||
|
refused_from 203.0.113.9 / || fail "the client of the probe was let through"
|
||||||
|
wait_for "the client's next request did not make its ban permanent" \
|
||||||
|
logged '"ban_expires":"permanent"'
|
||||||
|
echo "example-app: a probe for /.env bans the client, its next request for good"
|
||||||
|
|
||||||
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
|
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
|
||||||
fail "sv stop smallwebwaf failed"
|
fail "sv stop smallwebwaf failed"
|
||||||
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
|
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
|
||||||
@@ -83,6 +131,13 @@ main() {
|
|||||||
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
|
status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")"
|
||||||
[ "$status" = 0 ] || fail "docker stop left exit status $status"
|
[ "$status" = 0 ] || fail "docker stop left exit status $status"
|
||||||
echo "example-app: docker stop stops the container in order"
|
echo "example-app: docker stop stops the container in order"
|
||||||
|
|
||||||
|
docker rm "$CONTAINER" >/dev/null
|
||||||
|
start_container
|
||||||
|
refused || fail "the new container let the banned client through"
|
||||||
|
wait_for "smallwebwaf logged no request refused under the ban" \
|
||||||
|
logged '"action":"banned"'
|
||||||
|
echo "example-app: a new container on the same volume keeps the ban"
|
||||||
}
|
}
|
||||||
|
|
||||||
main "$@"
|
main "$@"
|
||||||
|
|||||||
+13
-1
@@ -1,6 +1,9 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# script/run: build bin/smallwebwaf with script/build and run it, with
|
# script/run: build bin/smallwebwaf with script/build and run it, with
|
||||||
# the settings in the environment.
|
# the settings in the environment. Unless SWWAF_STATE_DIR is set, the
|
||||||
|
# state files go in bin/state, beside the binary, and unless
|
||||||
|
# SWWAF_RULES_DIR is set, the rule files are those of share/rules.d,
|
||||||
|
# which the image ships.
|
||||||
set -eu
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -8,6 +11,15 @@ ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)"
|
|||||||
|
|
||||||
main() {
|
main() {
|
||||||
"$SCRIPT_DIR/build"
|
"$SCRIPT_DIR/build"
|
||||||
|
if [ -z "${SWWAF_STATE_DIR+set}" ]; then
|
||||||
|
SWWAF_STATE_DIR="$ROOT/bin/state"
|
||||||
|
export SWWAF_STATE_DIR
|
||||||
|
mkdir -p "$SWWAF_STATE_DIR"
|
||||||
|
fi
|
||||||
|
if [ -z "${SWWAF_RULES_DIR+set}" ]; then
|
||||||
|
SWWAF_RULES_DIR="$ROOT/share/rules.d"
|
||||||
|
export SWWAF_RULES_DIR
|
||||||
|
fi
|
||||||
exec "$ROOT/bin/smallwebwaf"
|
exec "$ROOT/bin/smallwebwaf"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# 00-default.rules: probes no real visitor sends, anchored at the site root
|
||||||
|
|
||||||
|
# id target action regex
|
||||||
|
env-file path ban (?i)^/\.env(\.[a-z]+)?$
|
||||||
|
vcs-dir path ban (?i)^/\.(git|svn|hg|bzr)(/|$)
|
||||||
|
secrets-dir path ban (?i)^/\.(aws|ssh|docker|kube)/
|
||||||
|
secret-file path ban (?i)^/\.(htpasswd|htaccess|npmrc|netrc|pgpass|git-credentials|bash_history|DS_Store)$
|
||||||
|
editor-dir path ban (?i)^/\.(vscode|idea)/
|
||||||
|
backup-file path ban (?i)^/[^/]+\.(php(\.[a-z0-9]+|~)|sql(\.[a-z0-9]+)?)$
|
||||||
|
log-file path ban (?i)^/(debug|error|access)\.log$
|
||||||
|
compose-file path ban (?i)^/(docker-)?compose\.ya?ml$
|
||||||
|
php-shell path ban (?i)^/(shell|c99|r57|wso|alfa)\.php$
|
||||||
|
scanner-agent user_agent ban (?i)\b(sqlmap|nikto|nuclei|masscan|zgrab|wpscan)\b
|
||||||
|
path-traversal uri block (\.\./){2,}
|
||||||
|
empty-agent user_agent log ^$
|
||||||
@@ -2,10 +2,14 @@
|
|||||||
set -euo pipefail
|
set -euo pipefail
|
||||||
|
|
||||||
# runit's run script for smallwebwaf, run again whenever smallwebwaf
|
# runit's run script for smallwebwaf, run again whenever smallwebwaf
|
||||||
# exits; the wait spaces out the restarts. exec, so that the signal
|
# exits; the wait spaces out the restarts. The state directory and every
|
||||||
# `sv stop` sends reaches smallwebwaf itself.
|
# file in it are given to the smallwebwaf user, so that a volume mounted
|
||||||
|
# there needs no change of owner; chown -R changes a symbolic link itself,
|
||||||
|
# never what it points to. exec, so that the signal `sv stop` sends
|
||||||
|
# reaches smallwebwaf itself.
|
||||||
main() {
|
main() {
|
||||||
sleep 1
|
sleep 1
|
||||||
|
chown -R smallwebwaf:smallwebwaf "${SWWAF_STATE_DIR:-/var/lib/smallwebwaf}"
|
||||||
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
|
exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user