Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b8cef7e88d | ||
|
|
6ec52e5b87 |
@@ -15,20 +15,23 @@ 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 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, 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
|
||||
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; 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, 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 [`SPEC.md`](SPEC.md). The survey of
|
||||
existing tools that led to the design is in [`EVALUATION.md`](EVALUATION.md).
|
||||
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).
|
||||
|
||||
## Getting started
|
||||
|
||||
@@ -100,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
|
||||
@@ -143,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
|
||||
|
||||
@@ -226,6 +232,23 @@ it, and the effective settings are logged at start.
|
||||
`********` 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
|
||||
@@ -235,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
|
||||
@@ -296,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
|
||||
@@ -335,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
|
||||
|
||||
@@ -376,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
|
||||
@@ -484,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
|
||||
@@ -678,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`.
|
||||
@@ -691,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
|
||||
|
||||
@@ -735,10 +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), 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).
|
||||
- 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
|
||||
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
+158
-2
@@ -4,6 +4,7 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -12,12 +13,15 @@ import (
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||
)
|
||||
|
||||
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
|
||||
@@ -115,6 +119,20 @@ type Config struct {
|
||||
// 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.
|
||||
@@ -166,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
|
||||
@@ -209,8 +234,16 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
|
||||
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",
|
||||
@@ -405,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) {
|
||||
@@ -709,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
|
||||
}
|
||||
|
||||
@@ -2,10 +2,13 @@ package config_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -47,8 +50,27 @@ const (
|
||||
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
||||
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
|
||||
metricsTopN = "SWWAF_METRICS_TOP_N"
|
||||
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"
|
||||
)
|
||||
|
||||
// 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"
|
||||
|
||||
@@ -198,6 +220,131 @@ func TestValuesAsSet(t *testing.T) {
|
||||
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
|
||||
}
|
||||
|
||||
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{
|
||||
logRemoteURL: "syslog+tls://logs.example:6514",
|
||||
logRemoteTLSCAFile: caFile,
|
||||
logRemoteBuffer: "500",
|
||||
logRemoteFacility: "daemon",
|
||||
logRemoteAppName: "fsn1app1/gitea",
|
||||
})
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -428,6 +575,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
||||
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
||||
}
|
||||
|
||||
hostname, _ := os.Hostname()
|
||||
|
||||
want := map[string]string{
|
||||
listenAddr: ":8080",
|
||||
upstreamURL: "http://127.0.0.1:8081",
|
||||
@@ -460,6 +609,11 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
|
||||
stateCounterInterval: "15m",
|
||||
metricsToken: "",
|
||||
metricsTopN: "50",
|
||||
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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -269,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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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(¶ms)
|
||||
fill(params)
|
||||
files := load(t, params)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
watch(t, files, lines)
|
||||
|
||||
// Each edit holds one entry, for a client the parts did not hold, and
|
||||
// takes the place of everything the part held.
|
||||
client := netip.MustParsePrefix("198.51.100.7/32")
|
||||
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "198.51.100.7/32", `+
|
||||
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
|
||||
wantTakenIn(t, lines, dir, bansJSON)
|
||||
wantEqual(t, bansJSON, params.Ledger.Snapshot(),
|
||||
[]bans.Ban{{Netblock: client, Start: midnight()}})
|
||||
|
||||
edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+
|
||||
`{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`)
|
||||
wantTakenIn(t, lines, dir, clientsJSON)
|
||||
wantEqual(t, clientsJSON, params.Limiter.Snapshot(),
|
||||
[]ratelimit.Client{{Client: client, History: ratelimit.History{Requests: 7}}})
|
||||
|
||||
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": [{"client": "198.51.100.7/32", `+
|
||||
`"country": "FR", "answered": "2026-10-06T00:00:00Z"}]}`)
|
||||
wantTakenIn(t, lines, dir, lookupsJSON)
|
||||
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
|
||||
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
|
||||
}
|
||||
|
||||
func TestOwnWritesAreNotTakenIn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
fill(params)
|
||||
files := load(t, params)
|
||||
watch(t, files, lines)
|
||||
|
||||
// Every file is written while watched, and then lookups.json edited:
|
||||
// the first edit taken in is that one.
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
edit(t, dir, lookupsJSON, `{"version": 1, "lookups": []}`)
|
||||
wantTakenIn(t, lines, dir, lookupsJSON)
|
||||
}
|
||||
|
||||
func TestFileRenamedOverAStateFileTakenIn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, bansJSON)
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
files := load(t, params)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
watch(t, files, lines)
|
||||
|
||||
// The admin mends bans.json.bad and moves it back, as editors that
|
||||
// save by renaming do with a file of their own: nothing is written
|
||||
// into bans.json itself. An edit of clients.json after it must be
|
||||
// taken in second.
|
||||
edit(t, dir, bansJSON+".bad", permanentBansJSON)
|
||||
|
||||
err = os.Rename(path+".bad", path)
|
||||
if err != nil {
|
||||
t.Fatalf("rename: %v", err)
|
||||
}
|
||||
|
||||
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
|
||||
|
||||
wantTakenIn(t, lines, dir, bansJSON)
|
||||
wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()})
|
||||
}
|
||||
|
||||
func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
watch(t, load(t, params), lines)
|
||||
|
||||
client := netip.MustParseAddr("203.0.113.9")
|
||||
|
||||
// An entry added, as an admin writes it, bans its netblock.
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", `+
|
||||
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
|
||||
wantTakenIn(t, lines, dir, bansJSON)
|
||||
|
||||
_, banned := params.Ledger.Check(client, midnight())
|
||||
if !banned {
|
||||
t.Error("the ban added to bans.json does not refuse")
|
||||
}
|
||||
|
||||
// The entry removed lifts the ban.
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
|
||||
wantTakenIn(t, lines, dir, bansJSON)
|
||||
|
||||
_, banned = params.Ledger.Check(client, midnight())
|
||||
if banned {
|
||||
t.Error("the ban removed from bans.json still refuses")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// It ends a ban's entry with a comma.
|
||||
const broken = "{\n \"version\": 1,\n \"bans\": [\n" +
|
||||
" {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n"
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, bansJSON)
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
params.Ledger.Load([]bans.Ban{permanentBan()})
|
||||
files := load(t, params)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
watch(t, files, lines)
|
||||
|
||||
// While smallwebwaf runs, the broken edit is left as it is: an edit
|
||||
// of clients.json, made after it and taken in, shows that it has been
|
||||
// seen.
|
||||
edit(t, dir, bansJSON, broken)
|
||||
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
|
||||
wantTakenIn(t, lines, dir, clientsJSON)
|
||||
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
|
||||
|
||||
// The next write sets it aside, logged with where the error is, and
|
||||
// writes bans.json again from what smallwebwaf still holds.
|
||||
err = files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
line := lines.waitFor(t, "set aside an edit of a state file that does not parse")
|
||||
message, _ := line["error"].(string)
|
||||
|
||||
if line["file"] != path+".bad" ||
|
||||
!strings.HasPrefix(message, path+", line 4, column 39: ") {
|
||||
t.Errorf("set aside with %v", line)
|
||||
}
|
||||
|
||||
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
|
||||
|
||||
if got := readFile(t, path+".bad"); got != broken {
|
||||
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
|
||||
}
|
||||
|
||||
if got := readFile(t, path); got != permanentBansJSON {
|
||||
t.Errorf("bans.json holds\n%s\nwant\n%s", got, permanentBansJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
files := load(t, params)
|
||||
|
||||
// One edit is taken in by the write of its file, before Watch runs,
|
||||
// and one by Watch.
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
watch(t, files, lines)
|
||||
edit(t, dir, bansJSON, permanentBansJSON)
|
||||
wantTakenIn(t, lines, dir, bansJSON)
|
||||
|
||||
wantMetric(t, scrape(t, params),
|
||||
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2)
|
||||
}
|
||||
|
||||
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
files := load(t, params)
|
||||
|
||||
// An edit taken in by Watch, which is then stopped.
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
stopped := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
files.Watch(ctx)
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
lines.waitFor(t, watching)
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
|
||||
byWatch := lines.waitFor(t, tookIn)
|
||||
|
||||
stop()
|
||||
<-stopped
|
||||
|
||||
// An edit taken in by the write of its file. Nothing logs after the
|
||||
// write, so the log is closed, and a write that does not log the edit
|
||||
// fails the test at once instead of waiting for the line.
|
||||
edit(t, dir, bansJSON, permanentBansJSON)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
close(lines)
|
||||
|
||||
byWrite := lines.waitFor(t, tookIn)
|
||||
|
||||
// The two lines differ only in their time.
|
||||
delete(byWatch, "time")
|
||||
delete(byWrite, "time")
|
||||
|
||||
if !maps.Equal(byWrite, byWatch) {
|
||||
t.Errorf("the write logged %v, where Watch logged %v", byWrite, byWatch)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
files := load(t, params)
|
||||
|
||||
edit(t, dir, bansJSON, `{"version": 1, "bans": [`)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
wantMetric(t, scrape(t, params),
|
||||
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1)
|
||||
}
|
||||
|
||||
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
lines := logInto(¶ms)
|
||||
files := load(t, params)
|
||||
|
||||
err := os.Remove(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("remove: %v", err)
|
||||
}
|
||||
|
||||
// Watch returns at once.
|
||||
files.Watch(t.Context())
|
||||
|
||||
line := lines.waitFor(t, "cannot watch the state files for edits")
|
||||
if line["level"] != "ERROR" {
|
||||
t.Errorf("logged as %v", line)
|
||||
}
|
||||
}
|
||||
|
||||
// midnight is the time of the tests' clock.
|
||||
@@ -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) {
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// The test is on write itself: a state file is read before it is
|
||||
// written, and a directory in its place fails that read first.
|
||||
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
|
||||
// A directory named bans.json cannot be renamed over.
|
||||
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700)
|
||||
if err != nil {
|
||||
t.Fatalf("mkdir: %v", err)
|
||||
}
|
||||
|
||||
err = write(dir, bansJSON, []byte("{}\n"))
|
||||
if err == nil {
|
||||
t.Error("writing over a directory did not fail")
|
||||
}
|
||||
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("read %s: %v", dir, err)
|
||||
}
|
||||
|
||||
if len(entries) != 1 || entries[0].Name() != bansJSON {
|
||||
t.Errorf("%s holds %v, want only bans.json", dir, entries)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user