Author SHA1 Message Date
clawbot b8cef7e88d Send every log line to a syslog server as well (closes #28)
check / check (push) Waiting to run
With SWWAF_LOG_REMOTE_URL set (syslog+udp, syslog+tcp or syslog+tls),
every line on stdout is also sent as the message of an RFC 5424 record,
octet-counted over TCP and TLS, from a bounded buffer that drops its
oldest line when full, so a slow or unreachable server holds up nothing.
Failed connections are retried with backoff; lines sent, dropped and
waiting are metrics. At a stop the lines still waiting get at most two
seconds. Standard library only: log/syslog writes only the older format.

Deviation: SWWAF_LOG_REMOTE_APP_NAME defaults to the host's name until
SWWAF_INSTANCE_NAME exists.

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

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

Model: opus-5-5
2026-10-06 14:18:13 +02:00
28 changed files with 2615 additions and 972 deletions
+131 -93
View File
@@ -13,21 +13,23 @@ JSON log line for every request.
Status: the first two milestones are built
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are seven parts of
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are six parts of
milestone 3: the static lists, the bans that broken rate limits lead to and the
JSON state files, which come next in the build order, `observe` mode and the
rest of the request log's fields, which come a little later, and the metrics
JSON state files with your edits taken in while it runs, which come next in the
build order, `observe` mode, which comes a little later, and the metrics
endpoint and the header size and the idle time as settings, which come last in
it. `smallwebwaf` passes each request to the app and the app's answer back,
unchanged, within its timeouts and size limits, works out each client's address,
bans a client that sends too many requests, refuses a client that comes from a
country you refuse or from a network you refuse, lets the networks you choose
through, keeps its bans, each client's counters and history, and GeoJS's answers
in JSON files across restarts, writes a JSON log line for every request, 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
it; and, from the stage after it, remote log sending. `smallwebwaf` passes each
request to the app and the app's answer back, unchanged, within its timeouts and
size limits, works out each client's address, bans a client that sends too many
requests, refuses a client that comes from a country you refuse or from a
network you refuse, lets the networks you choose through, keeps its bans, each
client's counters and history, and GeoJS's answers in JSON files across
restarts, takes in your edits of those 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).
@@ -66,9 +68,7 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
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
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. It
also gets the request's id in `X-Request-ID`, the same id as in the request's
log line (see `request_id` in "Request log" below).
`X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end.
- Enforces the timeouts and the size limits below. A limit passed before the
response has started gets `smallwebwaf`'s own answer: `408` for a client too
slow to send its request, `413` for a request body that is too large, `504`
@@ -103,9 +103,9 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
seen, how many of them the ban has refused, and how many bans the netblock had
before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent;
past that, the earliest ban of the netblock that has gone longest without a
request is dropped first. `bans.json` shows the bans and their notes, and a
restart lifts none (see "State files" below); lifting a ban by editing it
comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
request is dropped first. `bans.json` shows the bans and their notes, a
restart lifts none, and you add or lift a ban by editing it (see "State files"
below).
- 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
@@ -146,6 +146,9 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
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).
- 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
@@ -156,11 +159,6 @@ it, and the effective settings are logged at start.
- `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens.
- `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.
- `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).
@@ -228,17 +226,29 @@ it, and the effective settings are logged at start.
`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).
- `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.
- `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_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 the host's 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, a default that is not such a name stops the
start too.
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
@@ -248,8 +258,8 @@ ISO 3166-1 assigns today, and `xk` for Kosovo, in either case (`de` and `DE` are
the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn
`su`, stops the start, and so does a code on both country lists. `off` switches
a timeout, a size limit or a rate limit off;
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings, the state settings
and `SWWAF_METRICS_TOP_N` cannot be 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. At most 20,000 clients are kept,
with their counters and history, and an IPv6 client is counted by its /64. A new
@@ -262,42 +272,18 @@ GeoJS are kept, for 7 days each.
refused ones included:
```
{"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}
{"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}
```
A field that does not apply to a request is left out of its line, apart from
`type`, the fields from `time` to `user_agent`, `request_id`, `peer_ip`,
`client_group`, `country`, `action` and `duration_total`, which every line has.
- `time` is when the request arrived, in UTC. `instance` is
`SWWAF_INSTANCE_NAME`. `scheme` is the `X-Forwarded-Proto` a trusted proxy
sent, and otherwise `http`. `path` and `query` are as the client sent them.
- `request_id` is the `X-Request-ID` a trusted proxy sent, or a new random one
of 26 letters and digits when it sent none, or when the peer is not a trusted
proxy. A request passed to the app takes it there in `X-Request-ID`.
- `peer_ip` is the TCP peer, normally traefik. `forwarded_for` is the
`X-Forwarded-For` header as received, several lines of it joined with `, `.
`client_group` is the client as the rate limits count it: its IPv4 address as
a /32, or the /64 of its IPv6 address.
- `time` is when the request arrived, in UTC. `peer_ip` is the TCP peer,
normally traefik. `path` and `query` are as the client sent them.
- `country` is the client's country as GeoJS places it. 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
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.
- `action` is `forward` for a request passed to the app, `denied` for one
refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused
@@ -313,15 +299,6 @@ A field that does not apply to a request is left out of its line, apart from
`banned`, `country_denied` or `rate_limited`. `action` then names what was
done: `forward` for a request passed to the app, and another action, such as
`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`, and one that `SWWAF_DENY_NETS`, a ban or the
country lists refuse, or would refuse in `observe` mode. The byte totals come
with the byte limits.
- `limit_hit` is there for a request that broke a rate limit, and names the
window whose limit it went over: `minute`, `hour` or `day`, the shortest if it
went over several. `offence` is then `limit`.
@@ -329,19 +306,10 @@ A field that does not apply to a request is left out of its line, apart from
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.
- 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.
- `duration_total` and `duration_upstream_total` are in milliseconds.
No body is logged, and no header but those above. `smallwebwaf`'s own messages
(start, the settings, stop, errors) share the stream as JSON lines marked
No body and no other header is logged. `smallwebwaf`'s own messages (start, the
settings, stop, errors) share the stream as JSON lines marked
`"type":"process"`.
Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and
@@ -351,6 +319,30 @@ which it answers `431`, headers slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`,
whose connection it closes without an answer, and requests it cannot read at
all, which it answers itself, mostly with `400`.
### 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 is made again at once. A
failed attempt to connect is logged and followed by the next a second later,
twice as long after each further failure up to a minute, and a second again once
a connection is made. 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
@@ -390,9 +382,42 @@ without a field it needs, named with the entry's place in the file: a ban's
`netblock`, `start` or `expires`, which is `null` for a permanent ban; a
client's `client`, or the `start` of a window in which it has requests; an
answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or
`answered`. An edit made while `smallwebwaf` runs is overwritten by its next
write: taking it in comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
The AS number and AS name come with their lookup.
`answered`. The AS number and AS name come with their lookup.
While it runs, `smallwebwaf` watches `SWWAF_STATE_DIR` and takes in your edit of
a state file as soon as you save it: what the file then holds replaces what
`smallwebwaf` held for it, as if read at start. It tells its own writes from
yours by comparing the file with what it last read or wrote, and before it
writes a file it takes in any edit made since, so your edit is not overwritten;
a change `smallwebwaf` made after you opened the file, such as a new ban, is
lost when you save over it. An edit that would stop the start, because it does
not parse, has another `version` or leaves out a field an entry needs, does not
stop the running `smallwebwaf`: it keeps what it holds, and at the file's next
write renames your file to `<name>.bad`, such as `bans.json.bad`, writes the
file again from memory, and logs the file and where the error is. It waits for
that write because an editor's file can be read before the editor has finished
writing it. Mend the `.bad` file and move it back. A file you remove is written
again at its next write.
To ban a netblock, add an entry to `bans.json` with its `netblock`, its `start`
and its `expires`, `null` for a ban that never ends; its `notes` may be left
out. This `bans.json` bans `203.0.113.0/24` for good:
```json
{
"version": 1,
"bans": [
{
"netblock": "203.0.113.0/24",
"start": "2026-10-06T12:00:00Z",
"expires": null
}
]
}
```
To lift a ban, delete its entry. `smallwebwaf` then forgets the ban, so it does
not make the netblock's next ban longer.
## Metrics
@@ -431,7 +456,15 @@ other request. No metric carries a client's address.
- `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`.
`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
@@ -539,9 +572,9 @@ goes through the candidates one by one.
readable JSON files, written regularly and at every stop, so a restart loses
nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks
up the change. Nothing is read from disk while serving a request. The files
for the bans, the clients and the GeoJS answers are built (see "State files"
above); the others come with their features, and taking in an edit while
running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
for the bans, the clients and the GeoJS answers are built, with an edit taken
in while running (see "State files" above); the others come with their
features.
- Health checks, the metrics, and listing, adding and lifting bans or asking why
a given address was refused, all on the one port every request uses: under
`/_smallwebwaf/` on the app's own address, through traefik like any other
@@ -733,10 +766,14 @@ addresses are never sent to GeoJS.
answers.
- `internal/ratelimit`: the table of clients: counts each client's requests,
tells when one takes it over a rate limit, and keeps each client's history.
- `internal/state`: reads the state files at start, and writes them when they
are due and at the stop.
- `internal/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
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
installs Ubuntu's packages, nixpkgs, `runsvinit` and `smallwebwaf`, with
`share/smallwebwaf.run` as runit's `run` script for `smallwebwaf`.
@@ -746,8 +783,9 @@ addresses are never sent to GeoJS.
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
table of clients to 20,000, the GeoJS answers to 100,000 and the banned
netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and
`github.com/prometheus/client_golang` keeps the metrics and serves them. The
country codes are the list in `internal/config/config.go`.
`github.com/prometheus/client_golang` keeps the metrics and serves them, and
`github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file is saved.
The country codes are the list in `internal/config/config.go`.
## Entrypoints
@@ -790,9 +828,9 @@ so that they run in minimal containers.
## TODO
- The rest of milestone 3: taking in an admin's edits to the state files
(https://git.eeqj.de/sneak/smallwebwaf/issues/68) and exemptions; then the
rest of the design, in the order of the build order in [`SPEC.md`](SPEC.md).
- The rest of milestone 3: exemptions and the rest of the request log's fields;
then the rest of the design, in the order of the build order in
[`SPEC.md`](SPEC.md).
## Documents
+1
View File
@@ -3,6 +3,7 @@ module sneak.berlin/go/smallwebwaf
go 1.26.0
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
)
+2
View File
@@ -4,6 +4,8 @@ github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UF
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/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=
+50 -24
View File
@@ -26,9 +26,9 @@ const maxTextBytes = 256
type Rules struct {
// LimitBanDuration is how long a first ban lasts.
LimitBanDuration time.Duration
// LimitBanRepeatWindow is how soon after the netblock's last ban
// ended a broken limit counts as a repeat, which bans for
// repeatFactor times as long as that ban.
// LimitBanRepeatWindow is how soon after the end of the netblock's
// ban that ended last a broken limit counts as a repeat, which bans
// for repeatFactor times as long as that ban.
LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban; a ban that would be longer is
// permanent instead.
@@ -176,9 +176,23 @@ func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
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 last ban ended lasts
// LimitBanRepeatWindow after the netblock's ban that ended last lasts
// repeatFactor times as long as that one. A ban that would be longer
// 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
@@ -192,12 +206,22 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
bans, found := l.netblocks.Get(netblock)
if found {
last = &(*bans)[len(*bans)-1]
if last.ActiveAt(now) {
return *last
active := activeBan(*bans, now)
if active != nil {
return *active
}
notes.EarlierBans = last.Notes.EarlierBans + 1
// No ban is active, so each has an end. A ban an admin adds to
// bans.json can start after another and end before it, so the
// ban that ended last is looked for among them all.
ended := slices.MaxFunc(*bans, func(a, b Ban) int {
return a.Expires.Compare(b.Expires)
})
last = &ended
// The netblock's first ban held counts the bans it had before that
// one, since dropped to make room, and each ban held adds one.
notes.EarlierBans = (*bans)[0].Notes.EarlierBans + len(*bans)
}
notes.Request = notes.Request.cut()
@@ -282,21 +306,25 @@ func (l *Ledger) Snapshot() []Ban {
return held
}
// Load puts bans read from bans.json into a ledger that holds none yet,
// in the order they started, so that a netblock whose last ban started
// latest counts as the most recently seen. Each netblock is masked to its
// length, so that 203.0.113.9/24 is 203.0.113.0/24, and each text in the
// notes is cut to 256 bytes. Past MaxBans the earliest bans are dropped,
// as when they are made.
// Load puts bans read from bans.json 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. Each netblock is
// masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24, and
// each text in the notes is cut to 256 bytes. Past MaxBans the earliest
// bans are dropped, as when they are made.
func (l *Ledger) Load(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
bans = slices.Clone(bans)
slices.SortStableFunc(bans, func(a, b Ban) int {
return a.Start.Compare(b.Start)
})
l.mu.Lock()
defer l.mu.Unlock()
l.netblocks.Purge()
l.held = 0
l.v4Lengths, l.v6Lengths = nil, nil
for _, ban := range bans {
ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut()
@@ -318,11 +346,9 @@ func (l *Ledger) active(client netip.Addr, now time.Time) *Ban {
continue
}
// A ban is made only once the one before has ended, so only the
// last can be active.
last := &(*bans)[len(*bans)-1]
if last.ActiveAt(now) {
return last
ban := activeBan(*bans, now)
if ban != nil {
return ban
}
}
@@ -358,8 +384,8 @@ func (l *Ledger) add(ban Ban) {
}
// expiry returns when a ban for a broken limit made at now ends, or zero
// when it is permanent. last is the netblock's last ban, which has ended,
// or nil when it has none.
// when it is permanent. last is the netblock's ban that ended last, or nil
// when it has none.
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration
+104
View File
@@ -130,6 +130,75 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
}
}
func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
t.Parallel()
// As when an admin adds a permanent ban to bans.json with a start
// before that of the netblock's ban that has ended.
netblock := netip.MustParsePrefix("203.0.113.0/24")
permanent := bans.Ban{Netblock: netblock, Start: midnight().Add(-time.Hour)}
ended := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(time.Hour),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{permanent, ended})
now := midnight().Add(2 * time.Hour)
client := netip.MustParseAddr("203.0.113.9")
ban, banned := ledger.Find(client, now)
if !banned || !ban.Permanent() {
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
}
ban, banned = ledger.Check(client, now)
if !banned || !ban.Permanent() {
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
banned, ban)
}
// A limit broken now makes no shorter ban over the permanent one.
ban = ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
"want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
}
}
func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
t.Parallel()
// A 9-hour ban smallwebwaf made, the third in a row, and an admin's
// 1-hour ban added to bans.json over it, with no notes.
netblock := netip.MustParsePrefix("203.0.113.9/32")
nineHours := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(9 * time.Hour),
Notes: bans.Notes{EarlierBans: 2},
}
admins := bans.Ban{
Netblock: netblock,
Start: midnight().Add(time.Hour),
Expires: midnight().Add(2 * time.Hour),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{nineHours, admins})
// Once both have ended, a limit broken within the repeat window bans
// for three times the 9 hours, and the notes count the two bans
// before the 9-hour one, it, and the admin's.
ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
if ban.Expires.Sub(ban.Start) != 27*time.Hour || ban.Notes.EarlierBans != 4 {
t.Errorf("the next ban lasts %s with %d earlier bans, want 27h and 4",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
}
func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
t.Parallel()
@@ -151,6 +220,41 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
}
}
func TestLoadReplacesTheBansHeld(t *testing.T) {
t.Parallel()
// Room for three bans, so that the second load, were it added to the
// two bans held, would drop none of them to make room.
rules := defaultRules()
rules.MaxBans = 3
ledger := bans.New(rules)
kept := bans.Ban{Netblock: netip.MustParsePrefix("2001:db8::/64"), Start: midnight()}
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.0/24"), Start: midnight()},
kept,
})
// Loaded again without the first ban, as when an admin's edit of
// bans.json is taken in, that ban is lifted.
ledger.Load([]bans.Ban{kept})
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
if banned {
t.Error("a ban left out of the second load still refuses")
}
// The ledger holds one ban, so it makes two more without dropping any.
first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
bans.Notes{})
second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
bans.Notes{})
want := []bans.Ban{first, second, kept}
if got := ledger.Snapshot(); !slices.Equal(got, want) {
t.Errorf("the ledger holds %+v, want %+v", got, want)
}
}
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
t.Parallel()
+159 -28
View File
@@ -4,6 +4,7 @@
package config
import (
"crypto/x509"
"errors"
"fmt"
"log/slog"
@@ -19,6 +20,8 @@ import (
"strings"
"time"
"unicode/utf8"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
@@ -28,10 +31,6 @@ type Config struct {
ListenAddr string
// UpstreamURL is the app (SWWAF_UPSTREAM_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 or a rate limit would refuse is passed to the app instead, and
@@ -114,15 +113,26 @@ type Config struct {
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
// 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), 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
// log line at start.
@@ -174,8 +184,15 @@ var (
"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")
errShortToken = errors.New("is shorter than 32 characters")
errNotMode = errors.New("is not enforce or observe")
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")
)
// FromEnvironment reads the settings with lookupEnv, normally
@@ -183,11 +200,9 @@ var (
// that is set but invalid is an error that names it.
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
Observe: env.observe("SWWAF_MODE", "enforce"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
@@ -217,12 +232,18 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
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"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
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"),
}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME", hostname,
cfg.LogRemoteURL != nil)
for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) {
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
@@ -353,19 +374,6 @@ func (e *environment) countries(name, defaultValue string) []string {
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 := parseList(e.value(name, defaultValue))
e.check(name, err)
for i, header := range headers {
headers[i] = strings.ToLower(header)
}
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 {
@@ -430,6 +438,68 @@ func (e *environment) token(name string) string {
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.
func (e *environment) certificates(name string) *x509.CertPool {
path := e.value(name, "")
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. Its value is checked when it is set, and while lines
// are sent, when they would be sent with its default.
func (e *environment) appName(name, defaultValue string, sending bool) string {
_, set := e.lookupEnv(name)
value := e.value(name, defaultValue)
if (set || sending) && !isAppName(value) {
e.check(name, fmt.Errorf("%q %w", value, errNotAppName))
}
return value
}
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
// whole number of days such as 7d, or off.
func parseDuration(value string) (time.Duration, error) {
@@ -734,3 +804,64 @@ func parseUpstreamURL(value string) (*url.URL, error) {
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
}
+142 -27
View File
@@ -2,11 +2,13 @@ package config_test
import (
"bytes"
"crypto/x509"
"encoding/json"
"log/slog"
"maps"
"net/netip"
"os"
"path/filepath"
"slices"
"strings"
"testing"
@@ -48,13 +50,26 @@ const (
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"
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"
@@ -127,18 +142,6 @@ func TestDefaults(t *testing.T) {
wantNetblocks(t, cfg.DenyNets)
wantCountries(t, deniedCountries, cfg.DeniedCountries)
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)
}
}
func TestValuesAsSet(t *testing.T) {
@@ -217,18 +220,128 @@ func TestValuesAsSet(t *testing.T) {
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
}
func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) {
func TestRemoteLogSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
hostname, _ := os.Hostname()
if cfg.LogRemoteURL != nil || cfg.LogRemoteTLSCAs != nil ||
cfg.LogRemoteBuffer != 10000 || cfg.LogRemoteFacility != 16 ||
cfg.LogRemoteAppName != hostname {
t.Errorf("remote log settings %v, %v, %d, %d and %q, want no URL, no "+
"certificates, 10000, 16 and %q", cfg.LogRemoteURL, cfg.LogRemoteTLSCAs,
cfg.LogRemoteBuffer, cfg.LogRemoteFacility, cfg.LogRemoteAppName, hostname)
}
}
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{
instanceName: "fsn1app1/gitea",
logRequestHeaders: " Accept , X-Custom",
logRemoteURL: "syslog+tls://logs.example:6514",
logRemoteTLSCAFile: caFile,
logRemoteBuffer: "500",
logRemoteFacility: "daemon",
logRemoteAppName: "fsn1app1/gitea",
})
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)
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 != "fsn1app1/gitea" {
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)
}
}
@@ -395,7 +508,6 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
@@ -497,8 +609,11 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
stateCounterInterval: "15m",
metricsToken: "",
metricsTopN: "50",
instanceName: hostname,
logRequestHeaders: defaultLogRequestHeaders,
logRemoteURL: "",
logRemoteTLSCAFile: "",
logRemoteBuffer: "10000",
logRemoteFacility: "local0",
logRemoteAppName: hostname,
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
+7 -5
View File
@@ -197,19 +197,21 @@ func (g *GeoJS) Snapshot() []Answer {
return answers
}
// Load keeps answers read from lookups.json, in a GeoJS that keeps none
// yet, in the order they were last used, so that the one used longest
// 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) {
g.mu.Lock()
defer g.mu.Unlock()
answers = slices.Clone(answers)
slices.SortStableFunc(answers, func(a, b Answer) int {
return a.Used.Compare(b.Used)
})
g.mu.Lock()
defer g.mu.Unlock()
g.answers.Purge()
now := g.now()
for _, answer := range answers {
+48
View File
@@ -13,6 +13,7 @@ import (
"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"
)
@@ -44,6 +45,8 @@ type Metrics struct {
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.
@@ -105,6 +108,11 @@ func New(topN int) *Metrics {
"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{})
@@ -120,6 +128,7 @@ func New(topN int) *Metrics {
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize,
m.stateFileEditsTakenIn, m.stateFileEditsSetAside,
)
return m
@@ -165,6 +174,33 @@ func (m *Metrics) AddBansAndClients(
)
}
// 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)
@@ -230,6 +266,18 @@ func (m *Metrics) StateFileWritten(name string, size int, err error) {
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 {
+5 -8
View File
@@ -30,17 +30,14 @@ func (rq *request) banned(now time.Time) bool {
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.
// limitBroken counts the request for the rate limits at now, and reports
// whether it takes the client over one. 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
hit, over := rq.h.limiter.Count(group, now)
if !over {
return false
}
-28
View File
@@ -1,7 +1,6 @@
package proxy
import (
"crypto/rand"
"net/http"
"net/netip"
"slices"
@@ -49,33 +48,6 @@ func clientAddress(
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.
const ipv6GroupPrefix = 64
+13 -17
View File
@@ -14,14 +14,10 @@ const (
appHost = "app.example"
// client is the client's address, as a proxy names it.
client = "203.0.113.9"
// forwardedFor is the header that lists the client and its proxies,
// and forwardedProto the one that gives the scheme the client used.
forwardedFor = "X-Forwarded-For"
forwardedProto = "X-Forwarded-Proto"
// secure is the scheme a client reached traefik with, and plain the
// one smallwebwaf serves.
// forwardedFor is the header that lists the client and its proxies.
forwardedFor = "X-Forwarded-For"
// secure is the scheme a client reached traefik with.
secure = "https"
plain = "http"
)
// appHeaders is what the app tells about the headers it received.
@@ -69,13 +65,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
func clientAddressCases() []clientAddressCase {
trusted := map[string]string{trustedProxies: trustLocalhost}
forged := http.Header{
forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"},
forwardedProto: {secure},
"X-Real-Ip": {client},
forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
}
replaced := appHeaders{
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http",
}
return []clientAddressCase{{
@@ -91,10 +87,10 @@ func clientAddressCases() []clientAddressCase {
"outside the trusted proxies from the right",
env: trusted,
header: http.Header{
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost},
forwardedProto: {secure},
"X-Real-Ip": {client},
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
},
wantClient: client,
wantApp: appHeaders{
@@ -142,7 +138,7 @@ func requestWithHeaders(
Host: r.Host,
ForwardedFor: r.Header.Get(forwardedFor),
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
ForwardedProto: r.Header.Get(forwardedProto),
ForwardedProto: r.Header.Get("X-Forwarded-Proto"),
RealIP: r.Header.Get("X-Real-IP"),
})
})
+15 -29
View File
@@ -6,8 +6,6 @@ import (
"errors"
"io"
"net/http"
"os"
"reflect"
"slices"
"strings"
"sync/atomic"
@@ -16,7 +14,6 @@ import (
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
@@ -118,33 +115,27 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
}
}
// wantRequestFields checks the log line's fields about the request. Its
// time, its id and its timings are checked only for being there.
// wantRequestFields checks the log line's fields about the request.
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper()
hostname, _ := os.Hostname()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: hostname,
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent),
want := requestlog.Line{
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery,
Protocol: "HTTP/1.1", Status: http.StatusTeapot,
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent),
ResponseBytes: int64(received), UserAgent: "test-agent",
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal,
DurationUpstreamTotal: line.DurationUpstreamTotal,
}
if line.Line != want {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
_, err := time.Parse(time.RFC3339, line.Time)
if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
t.Errorf("log line has time %q, request_id %q and durations %v and %v",
line.Time, line.RequestID, line.DurationTotal, line.DurationUpstreamTotal)
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 {
t.Errorf("log line has time %q and durations %v and %v",
line.Time, line.DurationTotal, line.DurationUpstreamTotal)
}
}
@@ -380,13 +371,8 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
addr, out := startProxy(t, "http://"+localhost+":1", nil)
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
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")
wantLine(t, out.requestLine(t), http.StatusBadGateway,
requestlog.ActionUpstreamError)
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
return line["type"] == "process" && line["msg"] == "request to the app failed"
-2
View File
@@ -169,8 +169,6 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
defer rq.addToHistory()
refused := rq.check(r.Context())
rq.checked = time.Now()
if refused != nil {
rq.answer(*refused)
+1 -15
View File
@@ -35,10 +35,6 @@ const (
// localhost is where every test server listens, and so the address
// smallwebwaf sees each test's requests come from.
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.
@@ -71,8 +67,6 @@ const (
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
)
// output collects what smallwebwaf writes on stdout.
@@ -89,14 +83,6 @@ func (o *output) Write(p []byte) (int, error) {
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.
func (o *output) lines(t *testing.T) []map[string]any {
t.Helper()
@@ -136,7 +122,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
var found []logLine
for _, fields := range o.lines(t) {
if fields["type"] == requestType {
if fields["type"] == "request" {
found = append(found, decodeLine(t, fields))
}
}
+22 -113
View File
@@ -8,7 +8,6 @@ import (
"net/http/httputil"
"net/netip"
"os"
"strings"
"sync"
"sync/atomic"
"time"
@@ -47,9 +46,7 @@ type request struct {
peer netip.Addr
peerTrusted bool
start time.Time
// checked is when the checks were done, and upstreamStart when the
// request was handed to the app.
checked time.Time
// upstreamStart is when the request was handed to the app.
upstreamStart time.Time
// cancel ends the request to the app.
cancel context.CancelFunc
@@ -59,34 +56,26 @@ type request struct {
complete bool
// mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them, and notes the times
// below, from its own; once timersStopped is set, none of the timeouts
// acts any more.
// own, and the transport starts and stops them from its own; once
// timersStopped is set, none of them acts any more.
mu sync.Mutex
timersStopped bool
clientRequestTimer *time.Timer
upstreamRequestTimer *time.Timer
upstreamResponseTimer *time.Timer
// connected is when there was a connection to the app, requestSent
// 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
// requestSent is when the app had been sent the whole request.
requestSent time.Time
}
// newRequest starts handling r: it notes the time, counts the request as
// under way, works out the client, and starts the log line with what is
// known of the request.
// under way, and works out the client.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
h.metrics.RequestStarted()
start := time.Now()
peer := peerAddress(r)
trusted := h.config.TrustedProxies
peerTrusted := isInside(peer, trusted)
forwardedFor := r.Header.Values("X-Forwarded-For")
client := clientAddress(peer, forwardedFor, trusted)
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted)
rq := &request{
h: h,
@@ -95,37 +84,22 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
out: &responseWriter{ResponseWriter: w},
client: client,
peer: peer,
peerTrusted: peerTrusted,
peerTrusted: isInside(peer, trusted),
start: start,
line: requestlog.Line{
Time: requestlog.FormatTime(start),
Instance: h.config.InstanceName,
ClientIP: client.String(),
Method: r.Method,
Scheme: scheme(r, peerTrusted),
Host: r.Host,
Path: r.URL.EscapedPath(),
Query: r.URL.RawQuery,
Protocol: r.Proto,
Referer: r.Referer(),
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,
Time: requestlog.FormatTime(start),
ClientIP: client.String(),
PeerIP: peer.String(),
Method: r.Method,
Host: r.Host,
Path: r.URL.EscapedPath(),
Query: r.URL.RawQuery,
Protocol: r.Proto,
Referer: r.Referer(),
UserAgent: r.UserAgent(),
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 {
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
}
@@ -133,27 +107,6 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
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
// is known, before its body is read or anything reaches the app. It
// returns nil to let the request through. The checks of checkClient come
@@ -229,9 +182,7 @@ func (rq *request) forward(ctx context.Context) {
rq.cancel = cancel
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
GotConn: rq.gotConn,
WroteRequest: rq.wroteRequest,
GotFirstResponseByte: rq.gotFirstResponseByte,
WroteRequest: rq.wroteRequest,
})
out := rq.in.WithContext(ctx)
@@ -254,8 +205,7 @@ func (rq *request) forward(ctx context.Context) {
}
// rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set.
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set.
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme
@@ -264,7 +214,6 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
// the query as the client sent it.
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
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
@@ -278,7 +227,6 @@ func (rq *request) modifyResponse(res *http.Response) error {
// connection it takes over, not through rq.out.
rq.stopTimers()
rq.out.status = res.StatusCode
rq.line.Websocket = true
return nil
}
@@ -374,10 +322,6 @@ func (rq *request) finish() {
line := &rq.line
line.Status = rq.out.status
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 {
line.RequestBytes = rq.body.bytes.Load()
@@ -402,18 +346,12 @@ func (rq *request) finish() {
now := time.Now()
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() {
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()
line.DurationUpstreamTotal = requestlog.Milliseconds(upstreamDuration)
}
// Counted before the log line is written, so that the metrics count
@@ -426,17 +364,6 @@ func (rq *request) finish() {
}
}
// timing is the time from start to end in milliseconds, for one of the
// log line's timings, or nil when end is zero: what it times never
// happened.
func timing(start, end time.Time) *float64 {
if end.IsZero() {
return nil
}
return new(requestlog.Milliseconds(end.Sub(start)))
}
// addToHistory adds the request, which has ended, to its client's
// history.
func (rq *request) addToHistory() {
@@ -546,24 +473,6 @@ func (rq *request) bodyReceived() {
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:
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
-337
View File
@@ -1,337 +0,0 @@
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 whose length it does not announce 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\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, 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 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"])
}
}
+18 -28
View File
@@ -149,39 +149,26 @@ type Hit struct {
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
// not it is refused, and returns the client's requests in each window. It
// reports whether the request takes the client over a limit, and the
// window whose limit it goes over, the shortest if it is over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
// not it is refused. It reports whether the request takes the client over
// a limit, and the window whose limit it goes over, the shortest if it is
// over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
var (
requests [3]float64
hit Hit
)
var hit Hit
for i, b := range l.get(client).buckets() {
w := l.windows[i]
requests[i] = b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
requests := b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
}
}
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
return counts, hit, hit.Window != ""
return hit, hit.Window != ""
}
// Reset sets client's counts in every window back to zero. Its history
@@ -282,18 +269,21 @@ func (l *Limiter) Snapshot() []Client {
return clients
}
// Load puts clients read from clients.json into a table that holds none
// yet, in the order they were last seen, so that the least recently seen
// is dropped first. Buckets whose time has passed at now are emptied.
// 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) {
l.mu.Lock()
defer l.mu.Unlock()
clients = slices.Clone(clients)
slices.SortStableFunc(clients, func(a, b Client) int {
return a.History.LastSeen.Compare(b.History.LastSeen)
})
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
+3 -25
View File
@@ -62,14 +62,14 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight()
for range limit {
_, _, over := limiter.Count(client, start)
_, 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)
hit, over := limiter.Count(client, start)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want {
@@ -78,28 +78,6 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
}
}
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range 3 {
limiter.Count(client, start)
}
// A quarter into the next minute, the minute still covers three
// quarters of the bucket before, with its three requests, which count
// 2.25, and this one: 3.25. The hour and the day cover all four.
counts, _, _ := limiter.Count(client, start.Add(time.Minute+time.Minute/4))
want := ratelimit.Counts{Minute: 3.25, Hour: 4, Day: 4}
if counts != want {
t.Errorf("counts %+v, want %+v", counts, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel()
@@ -260,7 +238,7 @@ func wantCount(
) {
t.Helper()
_, hit, _ := limiter.Count(client, now)
hit, _ := limiter.Count(client, now)
if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want)
+296
View File
@@ -0,0 +1,296 @@
// 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"
"fmt"
"log/slog"
"net"
"net/url"
"os"
"strconv"
"sync/atomic"
"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, the next is made a second later,
// and retryDelayFactor times as long after each further failure in a
// row, up to a minute.
firstRetryDelay = time.Second
retryDelayFactor = 2
maxRetryDelay = 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, the record dropped, and
// a new one made at once. A failed attempt to connect is logged to
// processLog and followed by the next after firstRetryDelay,
// retryDelayFactor times as long after each further failure in a row up
// to maxRetryDelay. 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)
switch {
case ctx.Err() != nil:
return conn
case err != nil:
processLog.Warn("connecting 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)
continue
}
delay = firstRetryDelay
err = s.sendOn(ctx, conn)
if err == nil {
return conn
}
_ = conn.Close()
}
}
// 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.
func (s *Sender) write(conn net.Conn, record []byte) error {
_, err := conn.Write(record)
if err != nil {
s.dropped.Add(1)
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)
}
}
+458
View File
@@ -0,0 +1,458 @@
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"
"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.
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, loopback)
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, loopback)
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()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t, loopback)
addr := endpoint.Addr()
sender, logged, _ := run(t, params(remotelog.SchemeTCP, addr))
_, _ = sender.Write([]byte("one\n"))
synctest.Wait()
wantCounts(t, sender, 1, 0, 0)
conn := accept(t, endpoint)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "one"))
// The endpoint goes away. The sender notices when a record fails,
// and tries to connect again at once, then a second later, then two
// seconds after that.
_ = conn.Close()
_ = endpoint.Close()
writeUntilDropped(t, sender, 1)
sent := sender.Sent()
_, _ = sender.Write([]byte("two\n"))
time.Sleep(time.Second)
synctest.Wait()
endpoint = listen(t, addr.String())
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 = accept(t, endpoint)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "two"))
wantRetries(t, logged, "1s", "2s")
// Having connected, the sender waits a second again after the
// next failure.
_ = conn.Close()
_ = endpoint.Close()
writeUntilDropped(t, sender, 2)
wantRetries(t, logged, "1s", "2s", "1s")
})
}
func TestRecordsWaitingAtTheStopAreSent(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// Nothing listens at addr when the sender starts: it fails to
// connect, and waits a second to try again.
endpoint := listen(t, loopback)
addr := endpoint.Addr()
_ = endpoint.Close()
sender, logged, stop := run(t, params(remotelog.SchemeTCP, addr))
synctest.Wait()
wantRetries(t, logged, "1s")
_, _ = sender.Write([]byte("one\ntwo\n"))
endpoint = listen(t, addr.String())
// Stopped before that second is over, it connects to send them.
stop()
wantCounts(t, sender, 2, 0, 0)
frames := bufio.NewReader(accept(t, endpoint))
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 at addr, closed when the test ends.
func listen(t *testing.T, addr string) net.Listener {
t.Helper()
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", addr)
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
}
// 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()
}
}
// wantRetries checks that the sender logged a failed attempt to connect
// 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"] != "connecting 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.
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.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
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
}
+24 -72
View File
@@ -9,8 +9,6 @@ import (
"io"
"log/slog"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// The action a request line names: what smallwebwaf did with the
@@ -47,71 +45,32 @@ const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
// Line is one request's line in the request log. The field names, and
// 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.
// Line is one request's line in the request log. The field names are
// those of the "Request log" section of SPEC.md.
//
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Line struct {
Type string `json:"type"`
// The standard web log fields. Scheme is how the client reached
// smallwebwaf, or the trusted proxy in front of it.
Time string `json:"time"`
Instance string `json:"instance"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Scheme string `json:"scheme"`
Host string `json:"host"`
Path string `json:"path"`
Query string `json:"query"`
Protocol string `json:"protocol"`
Status int `json:"status"`
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
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"`
Type string `json:"type"`
Time string `json:"time"`
ClientIP string `json:"client_ip"`
PeerIP string `json:"peer_ip"`
Country string `json:"country"`
Method string `json:"method"`
Host string `json:"host"`
Path string `json:"path"`
Query string `json:"query"`
Protocol string `json:"protocol"`
Status int `json:"status"`
UpstreamStatus int `json:"upstream_status,omitempty"`
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
Referer string `json:"referer"`
UserAgent string `json:"user_agent"`
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 or ActionRateLimited.
WouldAction string `json:"would_action,omitempty"`
// Counts are the client's requests as the rate limits counted them
// with this one, for a request they counted.
Counts ratelimit.Counts `json:"counts,omitzero"`
// LimitHit is the window whose rate limit the request went over:
// minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"`
@@ -120,18 +79,11 @@ type Line struct {
// BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent".
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"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"`
// DurationTotal and DurationUpstreamTotal are in milliseconds.
DurationTotal float64 `json:"duration_total"`
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"`
}
// Write writes line to w as one JSON line marked "type":"request".
+1 -5
View File
@@ -50,11 +50,7 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
}
unset := []string{
"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",
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted",
"duration_upstream_total",
}
for _, name := range unset {
+74 -10
View File
@@ -18,6 +18,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
@@ -27,6 +28,11 @@ import (
// runit and docker wait a little longer before they kill the process.
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.
type Params struct {
// Version is the version of the binary, set when it is built.
@@ -69,16 +75,40 @@ func Run(ctx context.Context, params Params) int {
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 = remotelog.New(remotelog.Params{
URL: cfg.LogRemoteURL,
RootCAs: cfg.LogRemoteTLSCAs,
Buffer: cfg.LogRemoteBuffer,
Facility: cfg.LogRemoteFacility,
AppName: cfg.LogRemoteAppName,
})
stdout = io.MultiWriter(params.Stdout, remote)
processLog = requestlog.NewProcessLogger(stdout)
stopSending := startSending(ctx, remote, processLog)
defer stopSending()
}
// The state files give times in UTC.
now := func() time.Time { return time.Now().UTC() }
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
RequestLog: stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: now,
})
if remote != nil {
server.Metrics.AddRemoteLog(remote)
}
files, err := state.Load(state.Params{
Dir: cfg.StateDir,
@@ -113,9 +143,35 @@ func Run(ctx context.Context, params Params) int {
return serve(ctx, server.Server, listener, files, processLog)
}
// serve serves requests on listener, and writes the state files as they
// are due, until ctx is done. Then it gives the requests in progress
// shutdownTimeout to finish, and writes every state file.
// 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, and takes in an admin's edits of them, until ctx is done. Then it
// gives the requests in progress shutdownTimeout to finish, and writes
// every state file.
func serve(
ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, processLog *slog.Logger,
@@ -130,12 +186,18 @@ func serve(
defer stopWriting()
written := make(chan struct{})
watched := make(chan struct{})
go func() {
files.Run(writing)
close(written)
}()
go func() {
files.Watch(writing)
close(watched)
}()
select {
case err := <-served:
processLog.Error("serving failed", "error", err.Error())
@@ -165,13 +227,15 @@ func serve(
return 1
}
// Run's last write has ended, so nothing else writes the files. Every
// request has ended too, but for two kinds that Go's server does not
// wait for: one cut off because Shutdown timed out, and one whose
// connection switched protocols, such as a WebSocket. Such a request
// adds to its client's history only as it ends, which can be after
// this write, and then that request is missing from clients.json.
// 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
err = files.WriteAll()
if err != nil {
+255 -19
View File
@@ -10,6 +10,8 @@ import (
"net/http/httptest"
"os"
"path/filepath"
"slices"
"strconv"
"strings"
"sync"
"testing"
@@ -26,11 +28,14 @@ const (
// testVersion is the version the tests give smallwebwaf.
testVersion = "test"
// localhost is where the tests listen.
localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL"
stateDir = "SWWAF_STATE_DIR"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR"
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"
// greeting is what the tests' app answers.
greeting = "hello from the app"
)
@@ -217,8 +222,8 @@ func TestStateKeptAcrossRestarts(t *testing.T) {
rateLimitPerDay: "2",
// Neither comes due in the test: the files are written as
// smallwebwaf stops.
"SWWAF_STATE_WRITE_DELAY": "1h",
"SWWAF_STATE_COUNTER_INTERVAL": "1h",
stateWriteDelay: "1h",
stateCounterInterval: "1h",
}
// The two requests a day allows, and a stop.
@@ -247,12 +252,12 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
const scope = "SWWAF_BAN_SCOPE_V4_PREFIX"
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
"SWWAF_TRUSTED_PROXIES": localhost + "/32",
rateLimitPerDay: "1",
scope: "24",
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
trustedProxies: localhost + "/32",
rateLimitPerDay: "1",
scope: "24",
}
// 203.0.113.9's second request breaks the day limit, and bans
@@ -281,6 +286,142 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
})
}
func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) {
t.Parallel()
const (
// bans.json as an admin writes it with a ban, permanent, on
// 203.0.113.0/24, and with none.
oneBan = `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", ` +
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`
noBan = `{"version": 1, "bans": []}`
)
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
trustedProxies: localhost + "/32",
// No write comes due in the test, so only the watch on the
// directory can take the edits in.
stateWriteDelay: "1h",
stateCounterInterval: "1h",
}
runUntilStopped(t, env, func(url string) {
path := filepath.Join(dir, "bans.json")
saveUntilAnswered(t, path, oneBan, url, "203.0.113.9", http.StatusForbidden)
wantStatus(t, url, "198.51.100.7", http.StatusOK)
saveUntilAnswered(t, path, noBan, url, "203.0.113.9", http.StatusOK)
})
}
func 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(),
"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(),
"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()
@@ -384,9 +525,9 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
upstreamURL: appURL,
stateDir: dir,
"SWWAF_MODE": "enforce",
"SWWAF_STATE_WRITE_DELAY": "10s",
"SWWAF_STATE_COUNTER_INTERVAL": "15m",
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
"SWWAF_CLIENT_IDLE_TIMEOUT": "120s",
@@ -448,6 +589,69 @@ func wantGreeting(t *testing.T, url string) {
}
}
// 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) {
@@ -479,6 +683,40 @@ func wantRefused(t *testing.T, url string) {
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 {
@@ -497,7 +735,5 @@ func wantStatus(t *testing.T, url, from string, status int) {
_ = res.Body.Close()
if res.StatusCode != status {
t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status)
}
return res.StatusCode
}
+270 -76
View File
@@ -1,14 +1,17 @@
// Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and
// history, and lookups.json GeoJS's answers. Load reads them at start, and
// Run and WriteAll write them, each from a snapshot its part takes under
// its own lock, so that no request waits on the disk.
// 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"
@@ -17,8 +20,10 @@ import (
"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"
@@ -61,15 +66,26 @@ type Params struct {
// Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC.
Now func() time.Time
// ProcessLog receives what was read, and the writes that fail.
// ProcessLog 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.
// 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.
@@ -120,41 +136,28 @@ func Load(params Params) (*Files, error) {
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
}
var (
bansIn bansFile
clientsIn clientsFile
lookupsIn lookupsFile
)
f := &Files{params: params, sums: map[string][sha256.Size]byte{}}
err = errors.Join(
read(params.Dir, bansJSON, &bansIn),
read(params.Dir, clientsJSON, &clientsIn),
read(params.Dir, lookupsJSON, &lookupsIn),
)
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
}
held := make([]bans.Ban, 0, len(bansIn.Bans))
for _, entry := range bansIn.Bans {
held = append(held, entry.ban())
}
params.Ledger.Load(held)
params.Limiter.Load(clientsIn.Clients, params.Now())
params.GeoJS.Load(lookupsIn.Lookups)
params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", len(bansIn.Bans), "clients", len(clientsIn.Clients),
"lookups", len(lookupsIn.Lookups))
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
return &Files{params: params}, nil
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.
// 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()
@@ -172,7 +175,7 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue:
bansDue = nil
f.logFailure(f.writeBans())
f.logFailure(f.writeFile(bansJSON))
case <-interval.C:
f.logFailure(f.WriteAll())
}
@@ -182,7 +185,50 @@ func (f *Files) Run(ctx context.Context) {
// WriteAll writes every state file, as smallwebwaf stops. A file that
// fails does not keep the others from being written.
func (f *Files) WriteAll() error {
return errors.Join(f.writeBans(), f.writeClients(), f.writeLookups())
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON))
}
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
// state file as soon as it is saved: what the file holds replaces what
// smallwebwaf held for it. An edit that does not parse is left for the
// file's next write, which sets it aside, since a file can be read while
// an editor is still writing it. If Dir cannot be watched, that is
// logged, and an edit is taken in only before its file is written.
func (f *Files) Watch(ctx context.Context) {
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the state files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the state files for edits",
"directory", f.params.Dir)
for {
select {
case <-ctx.Done():
return
case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON:
f.fileChanged(name)
}
case err = <-watcher.Errors:
f.params.ProcessLog.Warn("watching the state files failed",
"error", err.Error())
}
}
}
// logFailure logs a write that failed.
@@ -193,52 +239,209 @@ func (f *Files) logFailure(err error) {
}
}
// writeBans writes bans.json.
func (f *Files) writeBans() error {
held := f.params.Ledger.Snapshot()
// 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()
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))}
for _, ban := range held {
file.Bans = append(file.Bans, newBanEntry(ban))
data, changed, err := f.readChanged(name)
if err != nil || !changed {
return
}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return fmt.Errorf("encode %s: %w", bansJSON, err)
}
return f.writeCounted(bansJSON, append(data, '\n'))
_ = f.takeInEdit(name, data)
}
// writeClients writes clients.json.
func (f *Files) writeClients() error {
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot())
// 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)
if err != nil {
return fmt.Errorf("encode %s: %w", clientsJSON, err)
return err
}
return f.writeCounted(clientsJSON, data)
// Counted before it is logged, so that the count is there once the
// log line is.
f.params.Metrics.StateFileEditTakenIn(name)
f.params.ProcessLog.Info("took in an edit of a state file",
"file", filepath.Join(f.params.Dir, name))
return nil
}
// writeLookups writes lookups.json.
func (f *Files) writeLookups() error {
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
if err != nil {
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
// 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.writeCounted(lookupsJSON, data)
return f.takeIn(name, data)
}
// writeCounted writes data to the state file name, as write does, and
// counts the write in the metrics.
func (f *Files) writeCounted(name string, data []byte) error {
err := write(f.params.Dir, name, data)
// 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. An error names the file and, where the
// JSON decoder tells it, the line and column, or else the entry.
func (f *Files) takeIn(name string, data []byte) (int, error) {
path := filepath.Join(f.params.Dir, name)
var entries int
switch name {
case bansJSON:
var file bansFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
held := make([]bans.Ban, 0, len(file.Bans))
for _, entry := range file.Bans {
held = append(held, entry.ban())
}
f.params.Ledger.Load(held)
entries = len(held)
case clientsJSON:
var file clientsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Limiter.Load(file.Clients, f.params.Now())
entries = len(file.Clients)
case lookupsJSON:
var file lookupsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups)
}
f.sums[name] = sha256.Sum256(data)
return entries, nil
}
// writeFile writes the state file name from what smallwebwaf holds. An
// edit made since smallwebwaf last read or wrote the file is taken in
// first, so that it is not overwritten, or set aside if it does not
// parse. A file that cannot be read, or an edit that cannot be set
// aside, is left as it is, and the write given up. Every write is counted
// in the metrics, and one that fails or is given up as a failure.
func (f *Files) writeFile(name string) error {
f.mu.Lock()
defer f.mu.Unlock()
data, changed, err := f.readChanged(name)
if err == nil && changed {
err = f.takeInEdit(name, data)
if err != nil {
err = f.setAside(name, err)
}
}
if err == nil {
data, err = f.encode(name)
if err != nil {
err = fmt.Errorf("encode %s: %w", name, err)
}
}
if err == nil {
err = write(f.params.Dir, name, data)
}
if err == nil {
// The file holds data from here on, even if the directory sync
// fails, so that its next read does not take it for an admin's
// edit.
f.sums[name] = sha256.Sum256(data)
err = syncDirectory(f.params.Dir)
}
f.params.Metrics.StateFileWritten(name, len(data), err)
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, Notes: ban.Notes}
@@ -389,28 +592,16 @@ func checkWritable(dir string) error {
return errors.Join(file.Close(), os.Remove(file.Name()))
}
// read reads the state file name in dir into file, a pointer to that
// file's struct, and checks its entries. A missing file leaves file as it
// is.
func read(dir, name string, file stateFile) error {
path := filepath.Join(dir, name)
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
if errors.Is(err, fs.ErrNotExist) {
return nil
}
if err != nil {
return err
}
// 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)
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)
@@ -463,7 +654,7 @@ func position(data []byte, err error) string {
// write writes data to the file name in dir so that a crash at any
// moment leaves either the old file or the new one, whole: data goes to a
// temporary file in the same directory, which is synced and renamed over
// name, and then the directory is synced, so that the rename lasts.
// 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"
@@ -475,10 +666,13 @@ func write(dir, name string, data []byte) error {
if err != nil {
_ = os.Remove(temporary)
return err
}
return err
}
// syncDirectory syncs dir to the disk, so that a rename in it lasts.
func syncDirectory(dir string) error {
directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
if err != nil {
return err
+481 -11
View File
@@ -3,7 +3,10 @@ package state_test
import (
"context"
"encoding/json"
"io/fs"
"log/slog"
"maps"
"net"
"net/http"
"net/http/httptest"
"net/netip"
@@ -28,6 +31,13 @@ const (
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
// What the process log says once Watch watches the directory, and as
// it takes in an edit.
watching = "watching the state files for edits"
tookIn = "took in an edit of a state file"
// maxLogLines is how many lines of the process log wait for a test to
// read them.
maxLogLines = 64
)
// permanentBansJSON is bans.json holding permanentBan.
@@ -290,7 +300,7 @@ func TestUnwritableDirectoryStopsTheStart(t *testing.T) {
}
}
// The two tests below run Run in a synctest bubble, where time is a clock
// The three tests below run Run in a synctest bubble, where time is a clock
// of the test's own: time.Sleep moves it on at once, and synctest.Wait
// returns once Run waits for its next write, so that every write due by
// then is on disk.
@@ -302,7 +312,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.WriteDelay = 10 * time.Second
run(t, load(t, params))
run(t, load(t, params).Run)
// A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it.
@@ -347,7 +357,7 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.CounterInterval = time.Minute
run(t, load(t, params))
run(t, load(t, params).Run)
// The files are removed once written, so that each interval shows
// them written again.
@@ -364,6 +374,44 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
})
}
func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.WriteDelay = 10 * time.Second
run(t, load(t, params).Run)
// A ban, and bans.json written with it.
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
midnight(), bans.Notes{})
time.Sleep(params.WriteDelay)
synctest.Wait()
// A second ban is to be written WriteDelay later. Just before
// then, an admin saves bans.json with the first ban lifted and
// another added.
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
midnight(), bans.Notes{})
time.Sleep(params.WriteDelay - time.Nanosecond)
synctest.Wait()
edit(t, dir, bansJSON, permanentBansJSON)
// The write takes the edit in first, and writes it back. The second
// ban, made after the admin opened the file, is lost, as "Edits
// while running" in SPEC.md says.
time.Sleep(time.Nanosecond)
synctest.Wait()
if got := readFile(t, filepath.Join(dir, bansJSON)); got != permanentBansJSON {
t.Errorf("bans.json holds\n%s\nwant the edit", got)
}
wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()})
})
}
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel()
@@ -459,24 +507,359 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) {
float64(len(permanentBansJSON)))
}
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Parallel()
dir := t.TempDir()
files := load(t, newParams(dir))
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
files := load(t, params)
// A directory named bans.json cannot be renamed over.
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700)
// bans.json is a socket, which cannot be opened as a file, even by
// root, as the tests run in Docker, but which a rename could replace.
// Whether it holds an edit cannot be told, so it is left as it is.
socket, err := (&net.ListenConfig{}).Listen(t.Context(), "unix", path)
if err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
_ = socket.Close()
}()
err = files.WriteAll()
if err == nil {
t.Error("writing with bans.json unreadable did not fail")
}
info, err := os.Lstat(path)
if err != nil || info.Mode().Type() != fs.ModeSocket {
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
}
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantWriteFailed(t, params, bansJSON)
}
func TestBrokenEditThatCannotBeSetAsideIsNotWrittenOver(t *testing.T) {
t.Parallel()
const broken = `{"version": 1, "bans": [`
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
files := load(t, params)
// A directory named bans.json.bad cannot be renamed over, so the
// broken edit cannot be set aside, and is left as it is.
edit(t, dir, bansJSON, broken)
err := os.Mkdir(path+".bad", 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
err = files.WriteAll()
if err == nil {
t.Error("writing over a directory did not fail")
t.Error("writing with bans.json.bad in the way did not fail")
}
if got := readFile(t, path); got != broken {
t.Errorf("bans.json holds\n%s\nwant the edit", got)
}
wantWriteFailed(t, params, bansJSON)
}
func TestEditOfEachFileTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// Each edit holds one entry, for a client the parts did not hold, and
// takes the place of everything the part held.
client := netip.MustParsePrefix("198.51.100.7/32")
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "198.51.100.7/32", `+
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON)
wantEqual(t, bansJSON, params.Ledger.Snapshot(),
[]bans.Ban{{Netblock: client, Start: midnight()}})
edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+
`{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`)
wantTakenIn(t, lines, dir, clientsJSON)
wantEqual(t, clientsJSON, params.Limiter.Snapshot(),
[]ratelimit.Client{{Client: client, History: ratelimit.History{Requests: 7}}})
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": [{"client": "198.51.100.7/32", `+
`"country": "FR", "answered": "2026-10-06T00:00:00Z"}]}`)
wantTakenIn(t, lines, dir, lookupsJSON)
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
}
func TestOwnWritesAreNotTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
fill(params)
files := load(t, params)
watch(t, files, lines)
// Every file is written while watched, and then lookups.json edited:
// the first edit taken in is that one.
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": []}`)
wantTakenIn(t, lines, dir, lookupsJSON)
}
func TestFileRenamedOverAStateFileTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// The admin mends bans.json.bad and moves it back, as editors that
// save by renaming do with a file of their own: nothing is written
// into bans.json itself. An edit of clients.json after it must be
// taken in second.
edit(t, dir, bansJSON+".bad", permanentBansJSON)
err = os.Rename(path+".bad", path)
if err != nil {
t.Fatalf("rename: %v", err)
}
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, bansJSON)
wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()})
}
func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
watch(t, load(t, params), lines)
client := netip.MustParseAddr("203.0.113.9")
// An entry added, as an admin writes it, bans its netblock.
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", `+
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned := params.Ledger.Check(client, midnight())
if !banned {
t.Error("the ban added to bans.json does not refuse")
}
// The entry removed lifts the ban.
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned = params.Ledger.Check(client, midnight())
if banned {
t.Error("the ban removed from bans.json still refuses")
}
}
func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Parallel()
// It ends a ban's entry with a comma.
const broken = "{\n \"version\": 1,\n \"bans\": [\n" +
" {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n"
dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
params := newParams(dir)
lines := logInto(&params)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
// While smallwebwaf runs, the broken edit is left as it is: an edit
// of clients.json, made after it and taken in, shows that it has been
// seen.
edit(t, dir, bansJSON, broken)
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, clientsJSON)
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
// The next write sets it aside, logged with where the error is, and
// writes bans.json again from what smallwebwaf still holds.
err = files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
line := lines.waitFor(t, "set aside an edit of a state file that does not parse")
message, _ := line["error"].(string)
if line["file"] != path+".bad" ||
!strings.HasPrefix(message, path+", line 4, column 39: ") {
t.Errorf("set aside with %v", line)
}
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
if got := readFile(t, path+".bad"); got != broken {
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
}
if got := readFile(t, path); got != permanentBansJSON {
t.Errorf("bans.json holds\n%s\nwant\n%s", got, permanentBansJSON)
}
}
func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
// One edit is taken in by the write of its file, before Watch runs,
// and one by Watch.
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
watch(t, files, lines)
edit(t, dir, bansJSON, permanentBansJSON)
wantTakenIn(t, lines, dir, bansJSON)
wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2)
}
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
// An edit taken in by Watch, which is then stopped.
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Watch(ctx)
close(stopped)
}()
lines.waitFor(t, watching)
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
byWatch := lines.waitFor(t, tookIn)
stop()
<-stopped
// An edit taken in by the write of its file. Nothing logs after the
// write, so the log is closed, and a write that does not log the edit
// fails the test at once instead of waiting for the line.
edit(t, dir, bansJSON, permanentBansJSON)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
close(lines)
byWrite := lines.waitFor(t, tookIn)
// The two lines differ only in their time.
delete(byWatch, "time")
delete(byWrite, "time")
if !maps.Equal(byWrite, byWatch) {
t.Errorf("the write logged %v, where Watch logged %v", byWrite, byWatch)
}
}
func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
files := load(t, params)
edit(t, dir, bansJSON, `{"version": 1, "bans": [`)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1)
}
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
err := os.Remove(dir)
if err != nil {
t.Fatalf("remove: %v", err)
}
// Watch returns at once.
files.Watch(t.Context())
line := lines.waitFor(t, "cannot watch the state files for edits")
if line["level"] != "ERROR" {
t.Errorf("logged as %v", line)
}
}
// midnight is the time of the tests' clock.
@@ -573,15 +956,16 @@ func load(t *testing.T, params state.Params) *state.Files {
return files
}
// run runs files' writes until the test ends.
func run(t *testing.T, files *state.Files) {
// run runs task, the Run or the Watch of state files, until the test
// ends.
func run(t *testing.T, task func(context.Context)) {
t.Helper()
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Run(ctx)
task(ctx)
close(stopped)
}()
@@ -591,6 +975,80 @@ func run(t *testing.T, files *state.Files) {
})
}
// watch runs files' Watch until the test ends, and waits until it
// watches the directory.
func watch(t *testing.T, files *state.Files, lines processLog) {
t.Helper()
run(t, files.Watch)
lines.waitFor(t, watching)
}
// processLog receives the lines of a process log, each a JSON object, for
// a test to wait for.
type processLog chan string
// logInto has params' process log write its lines into a new processLog,
// and returns that.
func logInto(params *state.Params) processLog {
lines := make(processLog, maxLogLines)
params.ProcessLog = slog.New(slog.NewJSONHandler(lines, nil))
return lines
}
// Write receives a line of the process log.
func (l processLog) Write(line []byte) (int, error) {
l <- string(line)
return len(line), nil
}
// waitFor returns the next line of the process log whose message is msg,
// passing over the lines before it, or nil if the log is closed first. It
// waits as long as that takes, so that a slow test process cannot fail
// the test.
func (l processLog) waitFor(t *testing.T, msg string) map[string]any {
t.Helper()
for line := range l {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil {
t.Fatalf("process log line %q is not JSON: %v", line, err)
}
if fields["msg"] == msg {
return fields
}
}
return nil
}
// wantTakenIn waits for the next edit taken in, and checks that it is of
// the state file name in dir.
func wantTakenIn(t *testing.T, lines processLog, dir, name string) {
t.Helper()
line := lines.waitFor(t, tookIn)
if line["file"] != filepath.Join(dir, name) {
t.Fatalf("took in %v, want an edit of %s", line, name)
}
}
// edit writes content to the state file name in dir, as an admin saves an
// edit of it.
func edit(t *testing.T, dir, name, content string) {
t.Helper()
err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
}
// wantEqual checks that the entries read back from file are those
// written.
func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
@@ -735,6 +1193,18 @@ func metric(t *testing.T, text, series string) float64 {
return 0
}
// wantWriteFailed checks that the metrics of params count one write of the
// state file name, and that it failed.
func wantWriteFailed(t *testing.T, params state.Params, name string) {
t.Helper()
got := scrape(t, params)
file := `{file="` + name + `"}`
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1)
}
// wantMetric checks the value of series in text, the metrics, as metric
// reads it.
func wantMetric(t *testing.T, text, series string, want float64) {
+35
View File
@@ -0,0 +1,35 @@
package state
import (
"os"
"path/filepath"
"testing"
)
// The test is on write itself: a state file is read before it is
// written, and a directory in its place fails that read first.
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
// A directory named bans.json cannot be renamed over.
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
err = write(dir, bansJSON, []byte("{}\n"))
if err == nil {
t.Error("writing over a directory did not fail")
}
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatalf("read %s: %v", dir, err)
}
if len(entries) != 1 || entries[0].Name() != bansJSON {
t.Errorf("%s holds %v, want only bans.json", dir, entries)
}
}