Compare commits

Author SHA1 Message Date
clawbot 7d49123874 Alerts to a JSON webhook, with a cooldown and an hourly summary (closes #26)
check / check (push) Waiting to run
SWWAF_ALERT_WEBHOOK_URL gets one JSON POST per alert, in SPEC.md's
schema, with SWWAF_ALERT_WEBHOOK_HEADERS: ban and permanent_ban, with
the ban's notes, in observe mode too, marked mode observe and worked
out only when the alert would be sent; source_failure for GeoJS;
file_error for a rule or state file with an error. SWWAF_ALERT_EVENTS
chooses; SWWAF_ALERT_COOLDOWN holds back repeats by netblock, file or
source; past SWWAF_ALERT_MAX_PER_HOUR the hour ends in one summary. A
bounded queue, retried with backoff, holds up no request; a 4xx other
than 408 and 429 gives the alert up. alerts.json keeps the queue, the
cooldowns and the hour. Nothing shows the URL's path or query.

Judgement call: the summary's event is summary, which SPEC.md omits.
Judgement call: an admin's ban raises no alert.

Model: opus-5-5
2026-10-07 01:36:59 +00:00
clawbot 5d6f6ffaf9 Admin endpoints for bans and clients on the single listener (closes #27)
check / check (push) Waiting to run
SWWAF_ADMIN_TOKEN, or its _FILE form, opens GET and POST
/_smallwebwaf/bans, DELETE /_smallwebwaf/bans/<client> and GET
/_smallwebwaf/clients/<ip>. Unset, they answer 404; a missing or wrong
token gets 401, in observe mode too. They go through every check, as
the metrics do. POST takes a netblock, not IPv4-mapped and without a
zone, or a client's address, a duration or permanent, and a reason, and
makes an admin ban even while another lasts. DELETE lifts every active
ban covering the address, kept and marked lifted. Bans come back as
bans.json entries; a client as clients.json holds it, with its bans.

Judgement call: answers leave out bans.json's version field.
Judgement call: DELETE takes an address, not a netblock.
Rule suppressed: gosec G304 on a test reading bans.json.

Model: opus-5-5
2026-10-07 01:13:16 +02:00
clawbot bff65f4e2f Settings given as files: the _FILE form of every setting (closes #87)
check / check (push) Waiting to run
Every setting X may instead be given as a file that X_FILE names, read
once at start: its contents, less one trailing newline, are the value,
checked as X would be. X and X_FILE both set, or a file that cannot be
read, stops the start with a message naming the variable. The logged
settings name the file, and mask a token read from one.
SWWAF_LOG_REMOTE_TLS_CA_FILE, whose value is a file already, has no
_FILE form. The health check reads only SWWAF_LISTEN_ADDR and
SWWAF_UPSTREAM_URL, so no other setting or file can fail it.

Judgement call: an invalid value read from a file is named as X, not X_FILE.
Rule suppressed: gosec G304 on reading the named file, as for the CA file.

Model: opus-5-5
2026-10-07 00:01:19 +02:00
clawbot ee9ba08a8a Bans an admin makes or lifts: the admin cause, a reason, lifted bans kept (closes #86)
check / check (push) Waiting to run
A bans.json entry without a cause gets the cause admin, written back so.
Bans whose cause is admin are never dropped and do not count toward
SWWAF_MAX_BANS, so setting a ban's cause to admin keeps it. Bans
smallwebwaf makes get a reason: the limit broken or the rule matched. A
lifted ban refuses nothing, is kept, and makes no later ban longer.
smallwebwaf_bans_made_total counts admin bans an edit adds while running;
earlier_bans counts admin in place of without_cause.

Judgement call: lifted lifts at once, whatever time it gives.
Judgement call: a lifted ban still counts in earlier_bans.
Known gap: a ban dropped from behind an admin's ban on its netblock leaves that netblock's later earlier_bans.

Model: opus-5-5
2026-10-06 23:09:52 +02:00
clawbot 0797e5def2 Send every log line to a syslog server as well (closes #28)
check / check (push) Waiting to run
With SWWAF_LOG_REMOTE_URL set (syslog+udp, syslog+tcp or syslog+tls),
every line on stdout is also sent as the message of an RFC 5424 record,
octet-counted over TCP and TLS, from a bounded buffer that drops its
oldest line when full, so a slow or unreachable server holds up nothing.
Failed connections are retried with backoff; lines sent, dropped and
waiting are metrics. At a stop the lines still waiting get at most two
seconds. SWWAF_LOG_REMOTE_APP_NAME defaults to SWWAF_INSTANCE_NAME; while
sending, an app name RFC 5424 does not allow stops the start. Standard
library only: log/syslog writes only the older format.

Model: opus-5-5
2026-10-06 22:02:36 +02:00
clawbot e77dfb6891 Rule files, and bans for a clear sign of attack (closes #24)
check / check (push) Waiting to run
Every *.rules file in SWWAF_RULES_DIR not named with a leading dot is
read at start, and again 2 seconds after the directory's last change.
Each request is checked against the rules after the rate limits: log
notes a match, block refuses with 403, ban refuses and bans the netblock
for SWWAF_ATTACK_BAN_DURATION, made permanent by its next request or
attack. path, query and uri are matched as the request line sent them;
header:Host and header:Transfer-Encoding are refused. Bans gain a cause.
The image ships 00-default.rules.

Judgement call: a header sent twice is matched with its values joined
by ", ".
Judgement call: SWWAF_MAX_BAN_DURATION does not cap a ban for an attack.
Not in this unit: offences for rule matches, with the error burst.

Model: opus-5-5
2026-10-06 20:38:35 +02:00
clawbot 74bdc6a449 Leave SWWAF_RATE_LIMIT_EXEMPT_PATHS out of the request rate limits (closes #77)
check / check (push) Waiting to run
A request is neither counted nor refused by the request rate limits
when its path as sent, the path the app receives, not percent-decoded,
starts with one of the comma-separated prefixes in
SWWAF_RATE_LIMIT_EXEMPT_PATHS, so /%61ssets/x is not under /assets/. A
request whose decoded path contains .. or a backslash, or whose path as
sent holds an encoded slash, is never exempt, since an app may act on
it as a path outside every prefix, such as /assets/..%2Flogin as
/login. The static lists, bans and the country lists still apply, and
its log line has no counts. The setting is empty by default, and a
prefix that does not start with / stops the start. README.md documents
it.

Model: opus-5-5
2026-10-06 18:47:13 +02:00
clawbot 808e69f442 Log the rest of the request log's fields (closes #79)
check / check (push) Successful in 3m23s
Each request log line now has the fields "Request log" in SPEC.md lists
whose features are built: instance (SWWAF_INSTANCE_NAME), scheme,
request_id (a trusted proxy's X-Request-ID or a new one, sent on to the
app), forwarded_for, client_group, content_type, content_length, the
headers SWWAF_LOG_REQUEST_HEADERS names, has_authorization, has_cookie,
websocket, response_content_type, cache_control, location, counts and
the timings. Authorization, Cookie and Set-Cookie values are never
logged. An entry of SWWAF_LOG_REQUEST_HEADERS that is not a header name,
or is Host or Transfer-Encoding, stops the start.

Deviation: counts has request totals only.
Deviation: SWWAF_INSTANCE_NAME is on request lines only.

Model: opus-5-5
2026-10-06 17:26:21 +02:00
52 changed files with 10810 additions and 643 deletions
+4
View File
@@ -167,6 +167,10 @@ RUN groupadd --system --gid 65532 smallwebwaf \
# smallwebwaf user at each start.
RUN mkdir /var/lib/smallwebwaf
# The default rule file, in SWWAF_RULES_DIR by default, where an app's
# Dockerfile can copy rule files of its own beside it.
COPY share/rules.d/00-default.rules /etc/smallwebwaf/rules.d/00-default.rules
# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv
# looks too.
COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run
+682 -138
View File
File diff suppressed because it is too large Load Diff
+3 -2
View File
@@ -32,7 +32,7 @@ from a directory of hand-editable text files.
- Defence against traffic floods that saturate the host's network link. That
needs help upstream of the host.
- A web UI or a configuration file. Settings are environment variables. Apart
from settings given as files (the `_FILE` form of any setting, such as
from settings given as files (the `_FILE` form of a setting, such as
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
files and the lookup database, the only files read are the rule files, which
hold one regex per line and nothing more elaborate.
@@ -298,7 +298,8 @@ it.
- A list set to an empty value is an empty list, and replaces the default.
- Every setting may instead be given as a file holding the value, named by the
setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for
secrets and long lists.
secrets and long lists. `SWWAF_LOG_REMOTE_TLS_CA_FILE`, whose value names a
file already, has no `_FILE` form.
- Settings, including those given as files, are read once at start; changing one
means restarting the container. The files `smallwebwaf` watches while it runs
are its state files, its rule files and the lookup database.
+609
View File
@@ -0,0 +1,609 @@
// Package alerts sends alerts on bans, on a source that fails and on a
// file with an error to the webhook SWWAF_ALERT_WEBHOOK_URL names, each
// as one JSON object, as the "Alert webhook schema" section of SPEC.md
// describes. A repeat within SWWAF_ALERT_COOLDOWN is held back, and so is
// an alert past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary. The
// others wait in a bounded queue, so that a slow or unreachable webhook
// never holds up a request. The state is written to alerts.json and read
// from it by the state package. Nothing logged names the webhook's URL,
// whose path or query can carry a secret.
package alerts
import (
"bytes"
"cmp"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net/http"
"net/netip"
"net/url"
"slices"
"sync"
"sync/atomic"
"time"
)
// The events an alert is for, as SWWAF_ALERT_EVENTS names them.
const (
// EventBan is a ban smallwebwaf made.
EventBan = "ban"
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
// made permanent.
EventPermanentBan = "permanent_ban"
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
// Core Rule Set, the anomaly thresholds and the reputation sources;
// nothing raises them yet.
EventWAFBlock = "waf_block"
EventAnomaly = "anomaly"
EventReputationHit = "reputation_hit"
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
EventSourceFailure = "source_failure"
// EventFileError is a rule file or state file edited while smallwebwaf
// runs that does not parse, or a state file that cannot be written.
EventFileError = "file_error"
// EventSummary is the summary of the alerts an hour held back past
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
EventSummary = "summary"
)
// Events returns every event SWWAF_ALERT_EVENTS can name, which is its
// default.
func Events() []string {
return []string{
EventBan, EventPermanentBan, EventWAFBlock, EventAnomaly,
EventReputationHit, EventSourceFailure, EventFileError,
}
}
const (
// queueSize is the most alerts that wait to be sent. Past it, the
// oldest is dropped.
queueSize = 1000
// sendTimeout bounds one request to the webhook.
sendTimeout = 10 * time.Second
// After a request to the webhook fails, the alert is sent again 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
// maxAnswerBytes is the most of the webhook's answer that is read.
maxAnswerBytes = 64 << 10
)
var (
errStatus = errors.New("the webhook answered")
// errRefused is a 4xx answer other than 408 and 429: the webhook
// refuses the alert itself, and would refuse it again.
errRefused = errors.New("the webhook refused the alert, answering")
)
// Params are what New needs.
type Params struct {
// WebhookURL is where each alert is posted (SWWAF_ALERT_WEBHOOK_URL),
// nil while it is unset and no alert is sent. WebhookHeaders are sent
// with each (SWWAF_ALERT_WEBHOOK_HEADERS).
WebhookURL *url.URL
WebhookHeaders http.Header
// Events are the events alerts are sent for (SWWAF_ALERT_EVENTS).
Events []string
// Cooldown is how long a repeat of an alert is held back
// (SWWAF_ALERT_COOLDOWN), 0 for no time. MaxPerHour is the most alerts
// sent in an hour (SWWAF_ALERT_MAX_PER_HOUR), 0 for no limit.
Cooldown time.Duration
MaxPerHour int
// Instance is SWWAF_INSTANCE_NAME, which every alert gives.
Instance string
// Now tells the time of an alert, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives the requests to the webhook that fail.
ProcessLog *slog.Logger
}
// Alert is one alert, as the webhook is sent it and alerts.json holds it,
// with the fields of the "Alert webhook schema" section of SPEC.md. ASN
// and ASName are empty until AS numbers are looked up.
//
//nolint:tagliatelle // SPEC.md's alert webhook schema names its fields in snake_case
type Alert struct {
Instance string `json:"instance"`
Time time.Time `json:"time"`
Event string `json:"event"`
Client netip.Addr `json:"client"`
Netblock netip.Prefix `json:"netblock"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"`
// Reason is a short sentence, and Detail what is particular to the
// event: for a file_error, its "file", and for a source_failure, its
// "source", which the cooldown tells repeats by.
Reason string `json:"reason"`
Detail map[string]any `json:"detail"`
// SuppressedRepeats is how many repeats of the alert the cooldown
// held back since the last one let through.
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Cooldown is, for an event on a netblock, or about a file or a source,
// when the last alert let through was raised, and how many repeats the
// cooldown has held back since, as alerts.json holds it.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Cooldown struct {
Event string `json:"event"`
Netblock netip.Prefix `json:"netblock"`
File string `json:"file,omitempty"`
Source string `json:"source,omitempty"`
Sent time.Time `json:"sent"`
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Hour is the hour under way, by the clock, as alerts.json holds it: when
// it started, how many alerts were let through in it, and how many were
// held back in it past MaxPerHour, by event, for its summary.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Hour struct {
Start time.Time `json:"start"`
Sent int `json:"sent"`
HeldBack map[string]int `json:"held_back"`
}
// State is what alerts.json holds: the cooldowns, the hour under way, and
// the alerts waiting to be sent, oldest first.
type State struct {
Cooldowns []Cooldown `json:"cooldowns"`
Hour Hour `json:"hour"`
Waiting []Alert `json:"waiting"`
}
// Queue takes the alerts raised, holds back those it must, and sends the
// others to the webhook. It is safe for concurrent use.
type Queue struct {
params Params
// httpClient follows no redirect: a redirect is a failure.
httpClient *http.Client
// queued receives a value when an alert joins the queue, unless one
// waits already, so that Run looks at the queue again.
queued chan struct{}
mu sync.Mutex
// cooldowns are the alerts last let through, by event and netblock,
// file or source.
cooldowns map[cooldownKey]*Cooldown
hour Hour
// waiting are the alerts waiting to be sent, oldest first.
waiting []*Alert
sent, failed, suppressed, dropped atomic.Int64
}
// cooldownKey is what makes an alert a repeat of another: the same event
// on the same netblock, and about the same file or source, as its detail
// names them. Each is empty for an alert without one.
type cooldownKey struct {
event string
netblock netip.Prefix
file string
source string
}
// cooldownKeyOf returns what makes another alert a repeat of alert.
func cooldownKeyOf(alert *Alert) cooldownKey {
file, _ := alert.Detail["file"].(string)
source, _ := alert.Detail["source"].(string)
return cooldownKey{alert.Event, alert.Netblock, file, source}
}
// New returns a Queue with no alert yet.
func New(params Params) *Queue {
return &Queue{
params: params,
httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
},
queued: make(chan struct{}, 1),
cooldowns: map[cooldownKey]*Cooldown{},
hour: Hour{HeldBack: map[string]int{}},
}
}
// Raise sends alert, which names its event and what is particular to it,
// unless no webhook is set or SWWAF_ALERT_EVENTS leaves its event out. It
// gives alert the instance and the time. An alert that repeats the last
// one let through less than Cooldown before is held back and counted,
// and the next one let through gives that count. Past MaxPerHour alerts
// let through in the hour under way, by the clock, an alert is held back
// for that hour's summary instead, which is sent once the hour has ended;
// it starts no cooldown, and the repeats held back before it are given by
// the next alert let through. Raise never waits: an alert let through
// joins the queue, from which Run sends it, and with queueSize alerts
// waiting the oldest is dropped.
func (q *Queue) Raise(alert Alert) {
if q.params.WebhookURL == nil || !slices.Contains(q.params.Events, alert.Event) {
return
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
alert.Instance = q.params.Instance
alert.Time = now
if q.repeat(&alert, now) {
q.suppressed.Add(1)
return
}
q.endHour(now)
if q.params.MaxPerHour > 0 && q.hour.Sent >= q.params.MaxPerHour {
q.hour.HeldBack[alert.Event]++
q.suppressed.Add(1)
return
}
q.startCooldown(&alert, now)
q.hour.Sent++
q.queue(&alert)
}
// WouldSend reports whether Raise would let an alert for event on
// netblock through now: a webhook is set, SWWAF_ALERT_EVENTS chooses
// event, no alert for event on netblock was let through less than
// Cooldown before, and fewer than MaxPerHour alerts have been let through
// in the hour under way. Unlike Raise, it counts nothing.
func (q *Queue) WouldSend(event string, netblock netip.Prefix) bool {
if q.params.WebhookURL == nil || !slices.Contains(q.params.Events, event) {
return false
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
last, found := q.cooldowns[cooldownKey{event: event, netblock: netblock}]
if q.params.Cooldown > 0 && found && now.Sub(last.Sent) < q.params.Cooldown {
return false
}
q.endHour(now)
return q.params.MaxPerHour == 0 || q.hour.Sent < q.params.MaxPerHour
}
// Run sends the alerts waiting, oldest first, until ctx is done. An alert
// stays in the queue until the webhook answers it with a 2xx status, or
// refuses it with a 4xx status other than 408 and 429: a refused alert is
// logged, counted as dropped, and given up, so that the next is sent. Any
// other request that fails is logged, and the alert sent again
// firstRetryDelay later, retryDelayFactor times as long after each
// further failure in a row, up to maxRetryDelay. Run also ends each hour
// as Raise does, so that the hour's summary is sent as it ends. With no
// webhook set, it returns at once.
func (q *Queue) Run(ctx context.Context) {
if q.params.WebhookURL == nil {
return
}
var (
retryDelay time.Duration
retryAt time.Time
)
for {
alert, untilHourEnds := q.next()
hourEnds := time.NewTimer(untilHourEnds)
var due <-chan time.Time // nil while no alert waits
if alert != nil {
due = time.After(time.Until(retryAt))
}
select {
case <-ctx.Done():
hourEnds.Stop()
return
case <-q.queued:
case <-hourEnds.C:
q.mu.Lock()
q.endHour(q.params.Now())
q.mu.Unlock()
case <-due:
err := q.send(ctx, alert)
switch {
case err == nil:
q.remove(alert)
q.sent.Add(1)
retryDelay = 0
retryAt = time.Time{}
case errors.Is(err, errRefused):
q.remove(alert)
q.failed.Add(1)
q.dropped.Add(1)
retryDelay = 0
retryAt = time.Time{}
q.params.ProcessLog.Warn("gave up an alert SWWAF_ALERT_WEBHOOK_URL refused",
"event", alert.Event, "error", err.Error())
case ctx.Err() == nil: // not cut off as smallwebwaf stops
q.failed.Add(1)
retryDelay = min(max(retryDelayFactor*retryDelay, firstRetryDelay),
maxRetryDelay)
retryAt = time.Now().Add(retryDelay)
q.params.ProcessLog.Warn("sending an alert to SWWAF_ALERT_WEBHOOK_URL failed",
"error", err.Error(), "sending_again_in", retryDelay.String())
}
}
hourEnds.Stop()
}
}
// Sent is how many alerts the webhook has taken.
func (q *Queue) Sent() int64 {
return q.sent.Load()
}
// Failed is how many requests to the webhook have failed.
func (q *Queue) Failed() int64 {
return q.failed.Load()
}
// Suppressed is how many alerts were held back: by the cooldown, and past
// MaxPerHour.
func (q *Queue) Suppressed() int64 {
return q.suppressed.Load()
}
// Dropped is how many alerts were dropped from a full queue, or given up
// as the webhook refused them.
func (q *Queue) Dropped() int64 {
return q.dropped.Load()
}
// Snapshot returns the queue's state, as alerts.json holds it, with the
// cooldowns sorted by netblock, then by event, file and source.
func (q *Queue) Snapshot() State {
q.mu.Lock()
defer q.mu.Unlock()
state := State{
Cooldowns: make([]Cooldown, 0, len(q.cooldowns)),
Hour: q.hour,
Waiting: make([]Alert, 0, len(q.waiting)),
}
state.Hour.HeldBack = maps.Clone(q.hour.HeldBack)
for _, cooldown := range q.cooldowns {
state.Cooldowns = append(state.Cooldowns, *cooldown)
}
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
})
for _, alert := range q.waiting {
state.Waiting = append(state.Waiting, *alert)
}
return state
}
// Load puts state, read from alerts.json, in place of the queue's state.
// Each cooldown's netblock is masked to its length, so that
// 203.0.113.9/24 is 203.0.113.0/24. Past queueSize alerts waiting, the
// oldest are dropped.
func (q *Queue) Load(state State) {
q.mu.Lock()
defer q.mu.Unlock()
q.cooldowns = map[cooldownKey]*Cooldown{}
for _, cooldown := range state.Cooldowns {
cooldown.Netblock = cooldown.Netblock.Masked()
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
q.cooldowns[key] = &cooldown
}
q.hour = state.Hour
q.hour.HeldBack = maps.Clone(state.Hour.HeldBack)
if q.hour.HeldBack == nil {
q.hour.HeldBack = map[string]int{}
}
q.waiting = nil
for _, alert := range state.Waiting {
q.queue(&alert)
}
}
// repeat reports whether alert, raised at now, repeats the last one let
// through less than Cooldown before, and counts it if it does.
func (q *Queue) repeat(alert *Alert, now time.Time) bool {
if q.params.Cooldown == 0 {
return false
}
last, found := q.cooldowns[cooldownKeyOf(alert)]
if !found || now.Sub(last.Sent) >= q.params.Cooldown {
return false
}
last.SuppressedRepeats++
return true
}
// startCooldown gives alert, let through at now, the count of the repeats
// held back since the last one let through, and notes alert as the last
// one let through.
func (q *Queue) startCooldown(alert *Alert, now time.Time) {
if q.params.Cooldown == 0 {
return
}
key := cooldownKeyOf(alert)
last, found := q.cooldowns[key]
if found {
alert.SuppressedRepeats = last.SuppressedRepeats
}
q.cooldowns[key] = &Cooldown{
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
Sent: now,
}
}
// endHour ends the hour under way, if now is past it: it queues that
// hour's summary when alerts were held back in it past MaxPerHour, and
// forgets the cooldowns that have run out with no repeat held back, which
// no alert needs any more.
func (q *Queue) endHour(now time.Time) {
start := now.Truncate(time.Hour)
if !start.After(q.hour.Start) {
return
}
heldBack := 0
for _, count := range q.hour.HeldBack {
heldBack += count
}
if heldBack > 0 {
q.queue(&Alert{
Instance: q.params.Instance,
Time: now,
Event: EventSummary,
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
Detail: map[string]any{
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
},
})
}
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
for key, cooldown := range q.cooldowns {
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
delete(q.cooldowns, key)
}
}
}
// queue adds alert to the alerts waiting, first dropping the oldest while
// queueSize wait, and has Run look at the queue again.
func (q *Queue) queue(alert *Alert) {
if len(q.waiting) == queueSize {
q.waiting = slices.Delete(q.waiting, 0, 1)
q.dropped.Add(1)
}
q.waiting = append(q.waiting, alert)
select {
case q.queued <- struct{}{}:
default: // a value waits already
}
}
// next returns the oldest alert waiting, nil when none waits, and how
// long it is until the hour under way ends.
func (q *Queue) next() (*Alert, time.Duration) {
q.mu.Lock()
defer q.mu.Unlock()
var oldest *Alert
if len(q.waiting) > 0 {
oldest = q.waiting[0]
}
return oldest, q.hour.Start.Add(time.Hour).Sub(q.params.Now())
}
// remove takes alert, which Run has sent or given up, out of the queue,
// unless it has been dropped from it, or Load has replaced the queue,
// since Run took it. Only the oldest alert is ever dropped, so alert is
// the oldest if it is there at all.
func (q *Queue) remove(alert *Alert) {
q.mu.Lock()
defer q.mu.Unlock()
if len(q.waiting) > 0 && q.waiting[0] == alert {
q.waiting = slices.Delete(q.waiting, 0, 1)
}
}
// send posts alert to the webhook as JSON, with WebhookHeaders, and
// returns an error unless the webhook answers with a 2xx status: one that
// wraps errRefused for a 4xx status other than 408 and 429. No error
// names the webhook's URL, whose path or query can carry a secret.
func (q *Queue) send(ctx context.Context, alert *Alert) error {
body, err := json.Marshal(alert)
if err != nil {
return fmt.Errorf("encode the alert: %w", err)
}
ctx, cancel := context.WithTimeout(ctx, sendTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
q.params.WebhookURL.String(), bytes.NewReader(body))
if err != nil {
return fmt.Errorf("make the request: %w", err)
}
maps.Copy(req.Header, q.params.WebhookHeaders)
req.Header.Set("Content-Type", "application/json")
res, err := q.httpClient.Do(req)
if err != nil {
// The client's error names the URL: only what went wrong is kept.
if urlErr, ok := errors.AsType[*url.Error](err); ok {
return urlErr.Err
}
return err
}
defer func() {
_ = res.Body.Close()
}()
// Read, so that the connection can be used again.
_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, maxAnswerBytes))
switch status := res.StatusCode; {
case status >= http.StatusOK && status < http.StatusMultipleChoices:
return nil
case status >= http.StatusBadRequest && status < http.StatusInternalServerError &&
status != http.StatusRequestTimeout && status != http.StatusTooManyRequests:
return fmt.Errorf("%w %s", errRefused, res.Status)
default:
return fmt.Errorf("%w %s", errStatus, res.Status)
}
}
+823
View File
@@ -0,0 +1,823 @@
package alerts_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// 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, the start
// of an hour: a wait lasts exactly as long as it should, however slowly
// the test process runs, and synctest.Wait returns once the queue has
// done all it can before time passes. The stand-in for the webhook
// answers without the network, since a request waiting on the network
// would keep that clock from moving on.
const (
// webhookURL is where the alerts are posted.
webhookURL = "https://alerts.example/smallwebwaf?team=ops"
// instance is the instance name every alert gives.
instance = "fsn1app1/gitea"
// started is when each test starts, as an alert gives it, and
// anHourOn an hour later.
started = "2000-01-01T00:00:00Z"
anHourOn = "2000-01-01T01:00:00Z"
// cooldown is the cooldown of most tests, the default.
cooldown = 15 * time.Minute
)
func TestAlertIsPostedAsJSONWithItsFieldsAndTheHeaders(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.WebhookHeaders = http.Header{
"Authorization": {"Bearer 0123456789abcdef"},
"X-Team": {"ops"},
}
webhook, q := start(t, params)
q.Raise(alerts.Alert{
Event: alerts.EventBan,
Client: netip.MustParseAddr("203.0.113.9"),
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Country: "DE",
Reason: "requests per minute over the limit of 1000",
Detail: map[string]any{"cause": "limit", "ban_expires": anHourOn},
})
synctest.Wait()
got := webhook.received()
if len(got) != 1 {
t.Fatalf("the webhook had %d requests, want 1", len(got))
}
if got[0].method != http.MethodPost || got[0].url != webhookURL {
t.Errorf("request %s %s, want POST %s", got[0].method, got[0].url, webhookURL)
}
for name, want := range map[string]string{
"Content-Type": "application/json",
"Authorization": "Bearer 0123456789abcdef",
"X-Team": "ops",
} {
if got[0].header.Get(name) != want {
t.Errorf("header %s is %q, want %q", name, got[0].header.Get(name), want)
}
}
wantAlert(t, got[0].alert, map[string]any{
"instance": instance,
"time": started,
"event": "ban",
"client": "203.0.113.9",
"netblock": "203.0.113.0/24",
"asn": "",
"as_name": "",
"country": "DE",
"reason": "requests per minute over the limit of 1000",
"detail": map[string]any{"cause": "limit", "ban_expires": anHourOn},
"suppressed_repeats": float64(0),
})
wantCounts(t, q, 1, 0, 0, 0)
})
}
func TestOnlyTheChosenEventsAreSent(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.Events = []string{alerts.EventSourceFailure, alerts.EventFileError}
webhook, q := start(t, params)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
q.Raise(alerts.Alert{Event: alerts.EventFileError, Reason: "a file error"})
q.Raise(alerts.Alert{Event: alerts.EventPermanentBan, Netblock: netblock(1)})
synctest.Wait()
wantEvents(t, webhook, alerts.EventFileError)
wantCounts(t, q, 1, 0, 0, 0)
})
}
func TestNothingIsQueuedWithoutAWebhook(t *testing.T) {
t.Parallel()
params := newParams()
params.WebhookURL = nil
q := alerts.New(params)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
if waiting := q.Snapshot().Waiting; len(waiting) != 0 {
t.Errorf("%d alerts wait, want none", len(waiting))
}
}
func TestRepeatWithinTheCooldownIsHeldBackAndCountedInTheNext(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
webhook, q := start(t, newParams())
raise := func(event string, n int) {
q.Raise(alerts.Alert{Event: event, Netblock: netblock(n)})
}
raise(alerts.EventBan, 1)
// The same event on the same netblock is a repeat; another netblock
// or another event is not.
time.Sleep(time.Minute)
raise(alerts.EventBan, 1)
raise(alerts.EventBan, 2)
raise(alerts.EventPermanentBan, 1)
time.Sleep(cooldown - time.Minute - time.Nanosecond)
raise(alerts.EventBan, 1)
// Once the cooldown has run out, the next one is sent with the
// count of those held back.
time.Sleep(time.Nanosecond)
raise(alerts.EventBan, 1)
// And starts the cooldown again.
time.Sleep(time.Minute)
raise(alerts.EventBan, 1)
synctest.Wait()
got := webhook.received()
wantEvents(t, webhook, alerts.EventBan, alerts.EventBan, alerts.EventPermanentBan,
alerts.EventBan)
for i, want := range []struct {
netblock int
repeats float64
}{{1, 0}, {2, 0}, {1, 0}, {1, 2}} {
alert := got[i].alert
if alert["netblock"] != netblock(want.netblock).String() ||
alert["suppressed_repeats"] != want.repeats {
t.Errorf("alert %d is for %v with %v repeats, want %s with %v", i,
alert["netblock"], alert["suppressed_repeats"], netblock(want.netblock),
want.repeats)
}
}
wantCounts(t, q, 4, 0, 3, 0)
})
}
func TestFileErrorAndSourceFailureRepeatOnlyForTheSameFileOrSource(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
webhook, q := start(t, params)
fileError := func(file string) alerts.Alert {
return alerts.Alert{
Event: alerts.EventFileError,
Detail: map[string]any{"file": file, "error": "line 2: an error"},
}
}
sourceFailure := func(source string) alerts.Alert {
return alerts.Alert{
Event: alerts.EventSourceFailure, Detail: map[string]any{"source": source},
}
}
// Another file, or another source, is no repeat.
q.Raise(fileError("/rules.d/50-a.rules"))
q.Raise(fileError("/rules.d/50-b.rules"))
q.Raise(fileError("/rules.d/50-a.rules"))
q.Raise(sourceFailure("geojs"))
q.Raise(sourceFailure("abuseipdb"))
q.Raise(sourceFailure("geojs"))
synctest.Wait()
// Each alert is named by its file, or its source.
got := make([]string, 0, len(webhook.received()))
for _, request := range webhook.received() {
detail, _ := request.alert["detail"].(map[string]any)
file, _ := detail["file"].(string)
source, _ := detail["source"].(string)
got = append(got, file+source)
}
want := []string{
"/rules.d/50-a.rules", "/rules.d/50-b.rules", "geojs", "abuseipdb",
}
if !slices.Equal(got, want) {
t.Errorf("the webhook was sent alerts for %v, want %v", got, want)
}
wantCounts(t, q, 4, 0, 2, 0)
// alerts.json keeps each file's cooldown: a new queue holds back
// the next for the first file, and sends the one for a third.
after := alerts.New(params)
after.Load(roundTrip(t, q.Snapshot()))
after.Raise(fileError("/rules.d/50-a.rules"))
after.Raise(fileError("/rules.d/50-c.rules"))
waiting := after.Snapshot().Waiting
if len(waiting) != 1 || waiting[0].Detail["file"] != "/rules.d/50-c.rules" {
t.Errorf("after loading, alerts wait %+v, want the one for 50-c.rules", waiting)
}
})
}
func TestNoCooldownSendsEveryRepeat(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.Cooldown = 0
webhook, q := start(t, params)
for range 3 {
q.Raise(alerts.Alert{Event: alerts.EventFileError})
time.Sleep(time.Minute)
}
synctest.Wait()
wantEvents(t, webhook, alerts.EventFileError, alerts.EventFileError,
alerts.EventFileError)
wantCounts(t, q, 3, 0, 0, 0)
})
}
func TestAlertsPastTheHourlyLimitAreRolledIntoOneSummary(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = 2
webhook, q := start(t, params)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(2)})
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(3)})
q.Raise(alerts.Alert{Event: alerts.EventPermanentBan, Netblock: netblock(4)})
q.Raise(alerts.Alert{Event: alerts.EventFileError})
// The summary is sent as the hour ends, and not before.
time.Sleep(time.Hour - time.Nanosecond)
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventBan)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventBan, alerts.EventSummary)
summary := webhook.received()[2].alert
wantAlert(t, summary, map[string]any{
"instance": instance,
"time": anHourOn,
"event": "summary",
"client": "",
"netblock": "",
"asn": "",
"as_name": "",
"country": "",
"reason": "3 alerts held back in the hour from 2000-01-01T00:00:00Z, " +
"past the 2 an hour SWWAF_ALERT_MAX_PER_HOUR allows",
"detail": map[string]any{
"hour": started,
"count": float64(3),
"events": map[string]any{
"ban": float64(1), "permanent_ban": float64(1), "file_error": float64(1),
},
},
"suppressed_repeats": float64(0),
})
// The next hour sends alerts again, and, with none held back, ends
// without a summary.
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(5)})
time.Sleep(time.Hour)
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventBan, alerts.EventSummary,
alerts.EventBan)
wantCounts(t, q, 4, 0, 3, 0)
})
}
func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = 1
webhook, q := start(t, params)
raise := func() {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
}
// The hour's one alert, and two repeats the cooldown holds back.
raise()
raise()
raise()
// Once the cooldown has run out, the next is past the hourly limit.
time.Sleep(cooldown)
raise()
// The next hour's first alert gives the two repeats, and the summary
// the alert past the limit.
time.Sleep(time.Hour - cooldown)
synctest.Wait()
raise()
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary, alerts.EventBan)
got := webhook.received()
if len(got) == 3 {
detail, _ := got[1].alert["detail"].(map[string]any)
repeats := got[2].alert["suppressed_repeats"]
if detail["count"] != float64(1) || repeats != float64(2) {
t.Errorf("the summary counts %v alerts, and the last alert gives %v "+
"repeats, want 1 and 2", detail["count"], repeats)
}
}
wantCounts(t, q, 3, 0, 3, 0)
})
}
func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
webhook, q := start(t, params)
webhook.set(failing)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
// A second after the first failure, then twice as long after each
// further one, up to a minute.
time.Sleep(200 * time.Second)
synctest.Wait()
after := make([]time.Duration, 0, len(webhook.received()))
for _, request := range webhook.received() {
after = append(after, request.at.Sub(midnight()))
}
want := []time.Duration{
0, time.Second, 3 * time.Second, 7 * time.Second, 15 * time.Second,
31 * time.Second, 63 * time.Second, 123 * time.Second, 183 * time.Second,
}
if !slices.Equal(after, want) {
t.Errorf("requests at %v, want %v", after, want)
}
wantCounts(t, q, 0, int64(len(want)), 0, 0)
if !strings.Contains(log.String(),
`"msg":"sending an alert to SWWAF_ALERT_WEBHOOK_URL failed"`) {
t.Errorf("process log %q names no failure", log.String())
}
// Once the webhook answers, the alert is sent, and leaves the
// queue.
webhook.set(answering)
time.Sleep(time.Minute)
synctest.Wait()
got := webhook.received()
if last := got[len(got)-1]; !last.answered ||
last.alert["netblock"] != netblock(1).String() {
t.Errorf("the last request was not the alert, answered")
}
wantCounts(t, q, 1, int64(len(want)), 0, 0)
if waiting := q.Snapshot().Waiting; len(waiting) != 0 {
t.Errorf("%d alerts still wait, want none", len(waiting))
}
})
}
func TestRefusedAlertIsGivenUpAndTheNextSent(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
webhook, q := start(t, params)
// 429 and 408 are failures, and the alert is sent again; 400 refuses
// it, and it is given up.
webhook.set(http.StatusTooManyRequests)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait()
webhook.set(http.StatusRequestTimeout)
time.Sleep(time.Second)
synctest.Wait()
webhook.set(refusing)
time.Sleep(2 * time.Second)
synctest.Wait()
// The next alert is sent at once.
webhook.set(answering)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(2)})
time.Sleep(time.Minute)
synctest.Wait()
// Each request, by when it was sent, and the netblock of its alert.
got := make([]string, 0, len(webhook.received()))
for _, request := range webhook.received() {
block, _ := request.alert["netblock"].(string)
got = append(got, request.at.Sub(midnight()).String()+" "+block)
}
want := []string{
"0s " + netblock(1).String(), "1s " + netblock(1).String(),
"3s " + netblock(1).String(), "3s " + netblock(2).String(),
}
if !slices.Equal(got, want) {
t.Errorf("requests %v, want %v", got, want)
}
wantCounts(t, q, 1, 3, 0, 1)
if !strings.Contains(log.String(),
`"msg":"gave up an alert SWWAF_ALERT_WEBHOOK_URL refused"`) {
t.Errorf("process log %q names no alert given up", log.String())
}
})
}
func TestFailedRequestIsLoggedWithoutTheURL(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
webhook, q := start(t, params)
webhook.set(hanging)
// The request is abandoned after 10 seconds, with an error from the
// HTTP client, which names the URL.
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
time.Sleep(11 * time.Second)
synctest.Wait()
logged := log.String()
if !strings.Contains(logged,
`"msg":"sending an alert to SWWAF_ALERT_WEBHOOK_URL failed"`) ||
strings.Contains(logged, "alerts.example") || strings.Contains(logged, "team=ops") {
t.Errorf("process log %q names no failure, or names the URL", logged)
}
})
}
func TestFullQueueDropsTheOldestAndRaiseNeverWaits(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = 0
webhook, q := start(t, params)
webhook.set(hanging)
// The webhook does not answer the first alert, while one more alert
// than the queue holds is raised: none waits, and the oldest, the
// one the webhook was sent, is dropped.
for n := range alerts.QueueSize + 1 {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(n)})
if n == 0 {
synctest.Wait()
}
}
if took := time.Since(midnight()); took != 0 {
t.Errorf("raising the alerts took %s, want no time", took)
}
wantCounts(t, q, 0, 0, 0, 1)
waiting := q.Snapshot().Waiting
if len(waiting) != alerts.QueueSize || waiting[0].Netblock != netblock(1) {
t.Fatalf("%d alerts wait, the first for %s, want %d, the first for %s",
len(waiting), waiting[0].Netblock, alerts.QueueSize, netblock(1))
}
// The request is abandoned after 10 seconds, and the webhook, which
// answers again, is sent the others, in order, a second later.
webhook.set(answering)
time.Sleep(11 * time.Second)
synctest.Wait()
got := webhook.received()
if len(got) != alerts.QueueSize+1 ||
got[0].alert["netblock"] != netblock(0).String() {
t.Fatalf("the webhook had %d requests, want %d, the first for %s",
len(got), alerts.QueueSize+1, netblock(0))
}
for i, request := range got[1:] {
if request.alert["netblock"] != netblock(i+1).String() {
t.Fatalf("request %d is for %v, want %s", i+1, request.alert["netblock"],
netblock(i+1))
}
}
wantCounts(t, q, alerts.QueueSize, 1, 0, 1)
})
}
func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = 1
before := alerts.New(params)
// Not sent: Run is not running. The repeat is held back by the
// cooldown, and the file error past the hourly limit.
before.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
before.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
before.Raise(alerts.Alert{Event: alerts.EventFileError})
time.Sleep(time.Minute)
webhook, after := start(t, params)
after.Load(roundTrip(t, before.Snapshot()))
// The new queue sends the alert waiting, holds back the repeat as
// the cooldown still runs, and sends the summary of the hour.
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan)
time.Sleep(time.Hour)
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary)
detail, _ := webhook.received()[1].alert["detail"].(map[string]any)
if detail["count"] != float64(1) {
t.Errorf("the summary counts %v alerts, want 1", detail["count"])
}
// The cooldown has run out, and the next one gives both repeats.
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait()
got := webhook.received()
if repeats := got[len(got)-1].alert["suppressed_repeats"]; repeats != float64(2) {
t.Errorf("the last alert gives %v repeats, want 2", repeats)
}
})
}
// How the stand-in for the webhook answers: with a status, or, hanging,
// not at all, until the request is abandoned.
const (
answering = http.StatusNoContent
failing = http.StatusServiceUnavailable
refusing = http.StatusBadRequest
hanging = 0
)
// standIn is a stand-in for the webhook. It notes each request it is
// sent.
type standIn struct {
mu sync.Mutex
answers int
requests []post
}
// post is a request the webhook was sent: when, its method, URL and
// headers, the alert it carried, and whether the webhook answered it with
// a 2xx status.
type post struct {
at time.Time
method string
url string
header http.Header
alert map[string]any
answered bool
}
// RoundTrip has the stand-in answer req, in place of the network. A
// request abandoned before the stand-in answers fails, as over the
// network.
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
answer := httptest.NewRecorder()
s.ServeHTTP(answer, req)
_ = req.Body.Close()
err := req.Context().Err()
if err != nil {
return nil, err
}
return answer.Result(), nil
}
// ServeHTTP notes the request, and answers it as the stand-in is set to.
func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
var alert map[string]any
_ = json.Unmarshal(body, &alert)
s.mu.Lock()
answers := s.answers
s.requests = append(s.requests, post{
at: time.Now(), method: r.Method, url: r.URL.String(), header: r.Header.Clone(),
alert: alert, answered: answers == answering,
})
s.mu.Unlock()
if answers == hanging {
<-r.Context().Done()
} else {
w.WriteHeader(answers)
}
}
// set sets how the stand-in answers: with the status answers, or hanging.
func (s *standIn) set(answers int) {
s.mu.Lock()
defer s.mu.Unlock()
s.answers = answers
}
// received returns the requests the stand-in has been sent so far.
func (s *standIn) received() []post {
s.mu.Lock()
defer s.mu.Unlock()
return slices.Clone(s.requests)
}
// lockedBuffer is a buffer the process log can write to while the test
// reads it.
type lockedBuffer struct {
mu sync.Mutex
buf bytes.Buffer
}
// Write adds p to the buffer.
func (b *lockedBuffer) Write(p []byte) (int, error) {
b.mu.Lock()
defer b.mu.Unlock()
return b.buf.Write(p)
}
// String returns what was written.
func (b *lockedBuffer) String() string {
b.mu.Lock()
defer b.mu.Unlock()
return b.buf.String()
}
// newParams returns the Params of most tests: the webhook at webhookURL,
// every event, the default cooldown and hourly limit, and the bubble's
// clock in UTC.
func newParams() alerts.Params {
webhook, err := url.Parse(webhookURL)
if err != nil {
panic(err)
}
return alerts.Params{
WebhookURL: webhook,
Events: alerts.Events(),
Cooldown: cooldown,
MaxPerHour: 60,
Instance: instance,
Now: func() time.Time { return time.Now().UTC() },
ProcessLog: slog.New(slog.DiscardHandler),
}
}
// start returns a stand-in for the webhook that answers, and a Queue that
// sends to it, run until the test ends.
func start(t *testing.T, params alerts.Params) (*standIn, *alerts.Queue) {
t.Helper()
webhook := &standIn{answers: answering}
q := alerts.New(params)
q.SetTransport(webhook)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
q.Run(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
return webhook, q
}
// midnight is when each test starts.
func midnight() time.Time {
return time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC)
}
// netblock returns the n-th netblock of a test, counted from 0.
func netblock(n int) netip.Prefix {
return netip.MustParsePrefix(fmt.Sprintf("203.0.%d.%d/32", 113+n/256, n%256))
}
// roundTrip returns state once written as JSON and read back, as
// alerts.json carries it from one start to the next.
func roundTrip(t *testing.T, state alerts.State) alerts.State {
t.Helper()
data, err := json.Marshal(state)
if err != nil {
t.Fatalf("encode: %v", err)
}
var read alerts.State
err = json.Unmarshal(data, &read)
if err != nil {
t.Fatalf("decode: %v", err)
}
return read
}
// wantAlert checks every field of an alert the webhook was sent.
func wantAlert(t *testing.T, got, want map[string]any) {
t.Helper()
if !reflect.DeepEqual(got, want) {
t.Errorf("alert %v, want %v", got, want)
}
}
// wantEvents checks the events of the alerts the webhook was sent, in
// order.
func wantEvents(t *testing.T, webhook *standIn, want ...string) {
t.Helper()
got := make([]string, 0, len(webhook.received()))
for _, request := range webhook.received() {
event, _ := request.alert["event"].(string)
got = append(got, event)
}
if !slices.Equal(got, want) {
t.Errorf("the webhook was sent %v, want %v", got, want)
}
}
// wantCounts checks the alerts q counts as sent, the requests it counts as
// failed, and the alerts it counts as held back and as dropped.
func wantCounts(
t *testing.T, q *alerts.Queue, sent, failed, suppressed, dropped int64,
) {
t.Helper()
if q.Sent() != sent || q.Failed() != failed || q.Suppressed() != suppressed ||
q.Dropped() != dropped {
t.Errorf("counts sent %d, failed %d, suppressed %d and dropped %d, "+
"want %d, %d, %d and %d", q.Sent(), q.Failed(), q.Suppressed(), q.Dropped(),
sent, failed, suppressed, dropped)
}
}
+12
View File
@@ -0,0 +1,12 @@
package alerts
import "net/http"
// QueueSize is the most alerts that wait to be sent.
const QueueSize = queueSize
// SetTransport has q's requests to the webhook go through transport
// instead of the network.
func (q *Queue) SetTransport(transport http.RoundTripper) {
q.httpClient.Transport = transport
}
+286
View File
@@ -0,0 +1,286 @@
package bans_test
import (
"net/netip"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
func TestBanWithoutACauseIsAnAdmins(t *testing.T) {
t.Parallel()
netblock := netip.MustParsePrefix("203.0.113.0/24")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{{Netblock: netblock, Start: midnight()}})
if got := ledger.Bans(netblock)[0].Cause; got != bans.CauseAdmin {
t.Errorf("the ban's cause is %q, want admin", got)
}
}
func TestAdminsBansAreNeverDroppedAndDoNotCountTowardMaxBans(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
adminsOnly := netip.MustParsePrefix("198.51.100.0/24")
both := netip.MustParsePrefix("203.0.113.1/32")
second := netip.MustParsePrefix("203.0.113.2/32")
third := netip.MustParsePrefix("203.0.113.3/32")
// Seen longest ago, a netblock with two of an admin's bans alone, and
// then one with an admin's ban before a ban smallwebwaf made: the one
// ban counted toward MaxBans.
ledger.Load([]bans.Ban{
{Netblock: adminsOnly, Start: midnight().Add(-3 * time.Hour), Cause: bans.CauseAdmin},
{Netblock: adminsOnly, Start: midnight().Add(-2 * time.Hour), Cause: bans.CauseAdmin},
{Netblock: both, Start: midnight().Add(-time.Hour), Cause: bans.CauseAdmin},
{
Netblock: both,
Start: midnight(),
Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
},
})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 2})
// A new ban drops the ban smallwebwaf made, and only that one.
ledger.BanForLimit(second, midnight(), bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 1})
if ledger.Bans(both)[0].Cause != bans.CauseAdmin {
t.Errorf("%s kept %+v, want the admin's ban", both, ledger.Bans(both))
}
// And the next drops that one.
ledger.BanForLimit(third, midnight(), bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 0, third: 1})
}
func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{Limit: 1000, Window: "minute"})
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
bans.Notes{RuleID: "git-dir", Target: "path"})
for _, tc := range []struct{ got, want string }{
{limit.Reason, "requests per minute over the limit of 1000"},
{attack.Reason, "matched the rule git-dir"},
} {
if tc.got != tc.want {
t.Errorf("the reason is %q, want %q", tc.got, tc.want)
}
}
}
func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
t.Parallel()
// An hour's ban lifted ten minutes after it started.
netblock := netip.MustParsePrefix("203.0.113.9/32")
lifted := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
Lifted: midnight().Add(10 * time.Minute),
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{lifted})
// While it would still last, it refuses nothing, and a limit broken
// bans for an hour, as a first broken limit does; the lifted ban is
// kept, and counted among the earlier bans.
now := midnight().Add(30 * time.Minute)
_, banned, _ := ledger.Check(netblock.Addr(), now)
if banned {
t.Error("the lifted ban refuses")
}
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
held := ledger.Bans(netblock)
if len(held) != 2 || held[0] != lifted {
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
}
}
func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
t.Parallel()
// A permanent ban for a clear sign of attack, lifted.
netblock := netip.MustParsePrefix("203.0.113.9/32")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{{
Netblock: netblock,
Start: midnight(),
Cause: bans.CauseAttack,
Lifted: midnight().Add(time.Hour),
}})
now := midnight().Add(2 * time.Hour)
_, banned, _ := ledger.Find(netblock.Addr(), now)
if banned {
t.Error("the lifted ban refuses")
}
active, permanent := ledger.Count(now)
if active != 0 || permanent != 0 {
t.Errorf("%d bans are active and %d permanent, want none", active, permanent)
}
// The next clear sign of attack bans for seven days, as a first does.
ban, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != 7*day {
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
}
}
func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
made, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{})
atStart := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
Start: midnight(),
}
// The bans read at the start were made before it.
ledger.Load([]bans.Ban{made, atStart})
if got := ledger.Made(bans.CauseAdmin); got != 0 {
t.Fatalf("%d bans made by an admin after the start's, want none", got)
}
// The admin keeps the ban smallwebwaf made, keeps the one read at the
// start, and adds one without a cause: that one alone is made.
kept := made
kept.Cause = bans.CauseAdmin
added := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.3/32"),
Start: midnight(),
}
ledger.LoadEdit([]bans.Ban{kept, atStart, added})
if ledger.Made(bans.CauseAdmin) != 1 || ledger.Made(bans.CauseLimit) != 1 {
t.Errorf("%d bans made by an admin and %d for a limit, want 1 of each",
ledger.Made(bans.CauseAdmin), ledger.Made(bans.CauseLimit))
}
}
func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
t.Parallel()
netblock := netip.MustParsePrefix("203.0.113.0/24")
ledger := bans.New(defaultRules())
// An hour's ban for a broken limit.
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// A minute later an admin bans the netblock for good, named by an
// address in it: that ban is made, and counts the other among the
// earlier bans.
now := midnight().Add(time.Minute)
want := bans.Ban{
Netblock: netblock,
Start: now,
Cause: bans.CauseAdmin,
Reason: "probes for logins",
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 1}},
}
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
"probes for logins")
if got != want {
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
}
wantChanged(t, ledger, true)
if made := ledger.Made(bans.CauseAdmin); made != 1 {
t.Errorf("%d bans made by an admin, want 1", made)
}
// It refuses once the ban for the limit has ended.
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
if !banned || ban != want {
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
ban, banned, want)
}
}
func TestLiftLiftsEveryActiveBanCoveringTheClient(t *testing.T) {
t.Parallel()
client := netip.MustParseAddr("203.0.113.9")
own := netip.MustParsePrefix("203.0.113.9/32")
wide := netip.MustParsePrefix("203.0.113.0/24")
other := netip.MustParsePrefix("203.0.113.10/32")
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{
// Ended an hour ago.
{
Netblock: own, Start: midnight().Add(-2 * time.Hour),
Expires: midnight().Add(-time.Hour), Cause: bans.CauseLimit,
},
// Active, on the client's address and on its /24.
{
Netblock: own, Start: midnight(), Expires: midnight().Add(time.Hour),
Cause: bans.CauseLimit,
},
{Netblock: wide, Start: midnight(), Cause: bans.CauseAdmin},
// Another client's.
{Netblock: other, Start: midnight(), Cause: bans.CauseAdmin},
})
now := midnight().Add(time.Minute)
lifted := ledger.Lift(client, now)
if len(lifted) != 2 || lifted[0].Lifted != now || lifted[1].Lifted != now {
t.Errorf("lifted %+v, want the two active bans covering the client", lifted)
}
wantChanged(t, ledger, true)
if _, banned, _ := ledger.Check(client, now); banned {
t.Error("the client is still banned")
}
if _, banned, _ := ledger.Check(other.Addr(), now); !banned {
t.Error("the other client's ban was lifted")
}
// The lifted bans are kept, and the one that had ended is not lifted.
covering := ledger.Covering(client)
if len(covering) != 3 || covering[0].Netblock != wide ||
!covering[1].Lifted.IsZero() || covering[2].Lifted != now {
t.Errorf("the bans covering the client are %+v, want the /24's and both "+
"of its own, the earlier not lifted", covering)
}
// With none active, nothing is lifted or changed.
if lifted = ledger.Lift(client, now); len(lifted) != 0 {
t.Errorf("lifted %+v again", lifted)
}
wantChanged(t, ledger, false)
}
+464 -114
View File
@@ -1,10 +1,13 @@
// Package bans is the ban ledger: the bans smallwebwaf makes on the
// netblocks of clients that break a rate limit, with their notes, as the
// "Bans" section of SPEC.md describes. The bans are kept in memory, and
// written to bans.json and read from it by the state package.
// netblocks of clients that break a rate limit or show a clear sign of
// attack, and those an admin makes, with their notes, as the "Bans"
// section of SPEC.md describes. The bans are kept in memory, and written
// to bans.json and read from it by the state package.
package bans
import (
"fmt"
"math"
"net/netip"
"slices"
"strings"
@@ -14,6 +17,17 @@ import (
"github.com/hashicorp/golang-lru/v2/simplelru"
)
// The causes of bans.
const (
// CauseLimit is a ban smallwebwaf made for a broken limit.
CauseLimit = "limit"
// CauseAttack is a ban smallwebwaf made for a clear sign of attack.
CauseAttack = "attack"
// CauseAdmin is a ban an admin made, or one smallwebwaf made that an
// admin keeps. It is never dropped.
CauseAdmin = "admin"
)
// repeatFactor is how many times as long as the netblock's last ban a ban
// for a limit broken again within the repeat window lasts.
const repeatFactor = 3
@@ -21,32 +35,44 @@ const repeatFactor = 3
// maxTextBytes is how much of each text in a ban's notes is kept.
const maxTextBytes = 256
// Rules are how long a ban for a broken limit lasts, and how many bans
// are held.
// Rules are how long a ban lasts, and how many bans are held.
type Rules struct {
// LimitBanDuration is how long a first ban lasts.
// LimitBanDuration is how long a first ban for a broken limit lasts.
LimitBanDuration time.Duration
// LimitBanRepeatWindow is how soon after the end of the netblock's
// ban that ended last a broken limit counts as a repeat, which bans
// for repeatFactor times as long as that ban.
// ban that ended last, other than one for a clear sign of attack, a
// broken limit counts as a repeat, which bans for repeatFactor times as
// long as that ban.
LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban; a ban that would be longer is
// permanent instead.
// MaxBanDuration is the longest ban for a broken limit; one that would
// be longer is permanent instead.
MaxBanDuration time.Duration
// MaxBans is the most bans held, at least one. Past it, the earliest
// ban of the netblock that has gone longest without a request is
// dropped.
// AttackBanDuration is how long a first ban for a clear sign of attack
// lasts.
AttackBanDuration time.Duration
// MaxBans is the most bans held whose cause is not CauseAdmin, at
// least one. Past it, the earliest such ban of the netblock that has
// gone longest without a request is dropped. Bans whose cause is
// CauseAdmin are held besides, and never dropped.
MaxBans int
}
// Ban is a ban on a netblock for a broken limit, the only kind of ban
// smallwebwaf makes so far.
// Ban is a ban on a netblock.
type Ban struct {
Netblock netip.Prefix
Start time.Time
// Expires is when the ban ends, zero for a permanent ban.
Expires time.Time
Notes Notes
// Cause is CauseLimit, CauseAttack or CauseAdmin.
Cause string
// Reason is a short text: for a ban smallwebwaf made, the limit broken
// or the rule that matched; for an admin's, what the admin wrote.
Reason string
// Lifted is when an admin lifted the ban, zero while no admin has. A
// lifted ban refuses nothing, and does not make the netblock's next
// ban longer.
Lifted time.Time
Notes Notes
}
// Permanent reports whether the ban never runs out.
@@ -54,9 +80,10 @@ func (b Ban) Permanent() bool {
return b.Expires.IsZero()
}
// ActiveAt reports whether the ban refuses requests at now.
// ActiveAt reports whether the ban refuses requests at now: it has not
// been lifted, and has not run out.
func (b Ban) ActiveAt(now time.Time) bool {
return b.Permanent() || now.Before(b.Expires)
return b.Lifted.IsZero() && (b.Permanent() || now.Before(b.Expires))
}
// Notes are what an admin needs to decide whether to lift a ban. The
@@ -66,23 +93,36 @@ func (b Ban) ActiveAt(now time.Time) bool {
type Notes struct {
// Country is the client's country, when it was looked up.
Country string `json:"country"`
// Limit, Window and Count are the limit that was broken, its window,
// "minute", "hour" or "day", and the count reached: the client's
// requests in the window, the one that broke the limit included.
// These are the requests that counted toward the ban, and the window
// is the time over which they came.
Limit int64 `json:"limit"`
Window string `json:"window"`
Count float64 `json:"count"`
// Request is the request that broke the limit.
// Limit, Window and Count are, for a ban for a broken limit, the limit
// that was broken, its window, "minute", "hour" or "day", and the
// count reached: the client's requests in the window, the one that
// broke the limit included. These are the requests that counted
// toward the ban, and the window is the time over which they came.
Limit int64 `json:"limit,omitempty"`
Window string `json:"window,omitempty"`
Count float64 `json:"count,omitempty"`
// RuleID and Target are, for a ban for a clear sign of attack, the id
// of the rule file rule that matched, and its target.
RuleID string `json:"rule_id,omitempty"`
Target string `json:"target,omitempty"`
// Request is the request that broke the limit, or that was the clear
// sign of attack.
Request Request `json:"request"`
// Requests is how many requests the netblock has sent since it was
// first seen, and Refused how many of them the ban has refused so
// far. Both go up with each request the ban refuses.
Requests int64 `json:"requests"`
Refused int64 `json:"refused"`
// EarlierBans is how many bans the netblock had before this one.
EarlierBans int `json:"earlier_bans"`
// EarlierBans is how many bans the netblock had before this one, by
// cause.
EarlierBans EarlierBans `json:"earlier_bans"`
}
// EarlierBans counts a netblock's bans before a ban, by cause.
type EarlierBans struct {
Limit int `json:"limit"`
Attack int `json:"attack"`
Admin int `json:"admin"`
}
// Request is a request in a ban's notes. Each text is cut to 256 bytes.
@@ -110,10 +150,13 @@ type Ledger struct {
// netblocks holds each banned netblock's bans, oldest first. Check and
// Find make each netblock they find the most recently seen.
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans.
// held is how many bans netblocks holds whose cause is not CauseAdmin,
// at most rules.MaxBans.
held int
// made is how many bans BanForLimit has made since the start.
made int
// made is how many bans have been made since the start, by cause: by
// the ledger, and by an admin, through BanForAdmin or in an edit of
// bans.json.
made map[string]int
// v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6
// netblocks that have been banned. Check looks for a ban at each of
// them, so that a ban read from bans.json refuses every client in its
@@ -124,9 +167,10 @@ type Ledger struct {
// New returns a Ledger with no ban yet.
func New(rules Rules) *Ledger {
// Every netblock held has a ban, so there are never more netblocks
// than rules.MaxBans, and the LRU never drops one itself.
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](rules.MaxBans, nil)
// The ledger drops bans itself, and never those whose cause is
// CauseAdmin, however many there are, so the LRU has no limit of its
// own: it keeps the netblocks in the order they were last seen.
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](math.MaxInt, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
@@ -135,45 +179,58 @@ func New(rules Rules) *Ledger {
rules: rules,
changed: make(chan struct{}, 1),
netblocks: netblocks,
made: map[string]int{},
}
}
// Changed receives a value after a ban is made, so that bans.json can be
// written. Several bans made before it is read leave one value.
// Changed receives a value after a ban is made, lifted or made permanent,
// so that bans.json can be written. Several changes before it is read
// leave one value.
func (l *Ledger) Changed() <-chan struct{} {
return l.changed
}
// Check is called for a request from client, at now. It reports whether
// a ban on a netblock client is in is active, and returns that ban, with
// the request counted among those it refused.
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
// the request counted among those it refused. A ban for a clear sign of
// attack is made permanent by the request: the netblock is malicious.
// The last result reports whether the request made the ban permanent.
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock()
defer l.mu.Unlock()
ban := l.active(client, now)
if ban == nil {
return Ban{}, false
return Ban{}, false, false
}
ban.Notes.Requests++
ban.Notes.Refused++
return *ban, true
madePermanent := ban.Cause == CauseAttack && !ban.Permanent()
if madePermanent {
ban.Expires = time.Time{}
l.markChanged()
}
return *ban, true, madePermanent
}
// Find is Check without counting the request among those the ban
// refused: in observe mode a ban refuses nothing.
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
// refused, and without making the ban permanent: in observe mode a ban
// refuses nothing. The last result reports whether Check would have made
// the ban permanent.
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock()
defer l.mu.Unlock()
ban := l.active(client, now)
if ban == nil {
return Ban{}, false
return Ban{}, false, false
}
return *ban, true
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent()
}
// activeBan returns the ban in bans, a netblock's bans oldest first, that
@@ -191,57 +248,155 @@ func activeBan(bans []Ban, now time.Time) *Ban {
}
// BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban. A first ban lasts LimitBanDuration. A ban made within
// LimitBanRepeatWindow after the netblock's ban that ended last lasts
// returns the ban, and true. A first ban lasts LimitBanDuration. A ban
// made within LimitBanRepeatWindow after the netblock's ban that ended
// last, other than one for a clear sign of attack or a lifted one, lasts
// repeatFactor times as long as that one. A ban that would be longer
// than MaxBanDuration is permanent instead. If a ban on netblock is still
// active, as when two of its requests break a limit at once, that ban is
// returned and no other is made. The ledger fills in the notes' Refused
// and EarlierBans itself.
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
// returned with false, and no other is made. The ledger fills in the
// notes' Refused and EarlierBans itself, and gives the ban the reason
// "requests per <Window> over the limit of <Limit>", from the notes.
func (l *Ledger) BanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
}
// WouldBanForLimit returns what BanForLimit would, without making the ban:
// what observe mode would have done.
func (l *Ledger) WouldBanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
}
// BanForAttack bans netblock at now for a clear sign of attack, with
// notes, and returns the ban, and whether it made it, as BanForLimit
// does. A first ban lasts AttackBanDuration; once the netblock has had
// one that was not lifted, the next is permanent. Its reason is "matched
// the rule <RuleID>".
func (l *Ledger) BanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true)
}
// WouldBanForAttack returns what BanForAttack would, without making the
// ban: what observe mode would have done.
func (l *Ledger) WouldBanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false)
}
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit
// or CauseAttack, made at now would be permanent, as BanForLimit or
// BanForAttack would make it. It works out nothing else of the ban.
func (l *Ledger) WouldBePermanent(
netblock netip.Prefix, now time.Time, cause string,
) bool {
l.mu.Lock()
defer l.mu.Unlock()
var last *Ban
bans, found := l.netblocks.Get(netblock)
if found {
active := activeBan(*bans, now)
if active != nil {
return *active
}
// 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)
var held []Ban
if bans, found := l.netblocks.Peek(netblock); found {
held = *bans
}
notes.Request = notes.Request.cut()
if cause == CauseAttack {
return l.attackExpiry(held, now).IsZero()
}
return l.limitExpiry(held, now).IsZero()
}
// limitReason is the reason of a ban for a broken limit, with notes.
func limitReason(notes Notes) string {
return fmt.Sprintf("requests per %s over the limit of %d", notes.Window, notes.Limit)
}
// attackReason is the reason of a ban for a clear sign of attack, with
// notes.
func attackReason(notes Notes) string {
return "matched the rule " + notes.RuleID
}
// BanForAdmin bans netblock at now for an admin, with reason, until
// expires, or for good when expires is zero, and returns the ban, whose
// cause is CauseAdmin. Unlike BanForLimit and BanForAttack, it makes the
// ban even while another on netblock is active, since the admin asked
// for this one. The ledger fills in the notes' EarlierBans, and counts
// the ban among those made.
func (l *Ledger) BanForAdmin(
netblock netip.Prefix, now, expires time.Time, reason string,
) Ban {
l.mu.Lock()
defer l.mu.Unlock()
ban := Ban{
Netblock: netblock,
Start: now,
Expires: l.expiry(last, now),
Notes: notes,
Netblock: netblock.Masked(), Start: now, Expires: expires, Cause: CauseAdmin,
Reason: reason,
}
l.add(ban)
l.made++
select {
case l.changed <- struct{}{}:
default: // a value is waiting already
held, found := l.netblocks.Get(ban.Netblock)
if found {
ban.Notes.EarlierBans = earlierBans(*held)
}
l.add(ban)
l.made[CauseAdmin]++
l.markChanged()
return ban
}
// Lift lifts, at now, every ban active then on a netblock client is in,
// as an admin does, and returns those bans. A lifted ban is kept, refuses
// nothing, and does not make the netblock's next ban longer.
func (l *Ledger) Lift(client netip.Addr, now time.Time) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
var lifted []Ban
for _, bans := range l.covering(client) {
for i := range *bans {
ban := &(*bans)[i]
if ban.ActiveAt(now) {
ban.Lifted = now
lifted = append(lifted, *ban)
}
}
}
if len(lifted) > 0 {
l.markChanged()
}
return lifted
}
// Covering returns every ban held on a netblock client is in, active or
// not, sorted by netblock, and each netblock's bans oldest first. It is
// not a request from client, and leaves when the netblocks were last seen
// unchanged.
func (l *Ledger) Covering(client netip.Addr) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
var held []Ban
for _, bans := range l.covering(client) {
held = append(held, *bans...)
}
slices.SortStableFunc(held, func(a, b Ban) int {
return a.Netblock.Compare(b.Netblock)
})
return held
}
// Bans returns the bans held on netblock, oldest first. It is not a
// request from netblock, and leaves when it was last seen unchanged.
func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
@@ -256,17 +411,19 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans)
}
// Made returns how many bans the ledger has made since the start; bans
// read from bans.json are not among them.
func (l *Ledger) Made() int {
// Made returns how many bans for cause have been made since the start:
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
// them. The bans read from bans.json at the start are not among them.
func (l *Ledger) Made(cause string) int {
l.mu.Lock()
defer l.mu.Unlock()
return l.made
return l.made[cause]
}
// Count returns how many of the bans held are active at now, and how many
// are permanent.
// of those are permanent. A lifted ban is neither.
func (l *Ledger) Count(now time.Time) (int, int) {
l.mu.Lock()
defer l.mu.Unlock()
@@ -275,10 +432,12 @@ func (l *Ledger) Count(now time.Time) (int, int) {
for _, bans := range l.netblocks.Values() {
for _, ban := range *bans {
if ban.ActiveAt(now) {
active++
if !ban.ActiveAt(now) {
continue
}
active++
if ban.Permanent() {
permanent++
}
@@ -306,30 +465,148 @@ func (l *Ledger) Snapshot() []Ban {
return held
}
// 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.
// Load puts bans read from bans.json at the start into the ledger, in
// place of the bans it holds, in the order they started, so that a
// netblock whose last ban started latest counts as the most recently
// seen. A ban without a cause is an admin's, and gets CauseAdmin. Each
// netblock is masked to its length, so that 203.0.113.9/24 is
// 203.0.113.0/24, and each text in the notes is cut to 256 bytes. Past
// MaxBans the earliest bans whose cause is not CauseAdmin are dropped, as
// when they are made.
func (l *Ledger) Load(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
l.load(bans)
}
// LoadEdit is Load for an admin's edit of bans.json, taken in while
// smallwebwaf runs. Each ban in it whose cause is CauseAdmin, and which
// the ledger did not hold, with the same netblock and start, is one the
// admin made, and is counted among the bans made.
func (l *Ledger) LoadEdit(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
l.made[CauseAdmin] += l.load(bans)
}
// load does what Load describes, and returns how many of bans are bans
// whose cause is CauseAdmin that the ledger did not hold before.
func (l *Ledger) load(bans []Ban) int {
bans = slices.Clone(bans)
added := 0
for i := range bans {
ban := &bans[i]
ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut()
if ban.Cause == "" {
ban.Cause = CauseAdmin
}
if ban.Cause == CauseAdmin && !l.holds(ban.Netblock, ban.Start) {
added++
}
}
slices.SortStableFunc(bans, func(a, b Ban) int {
return a.Start.Compare(b.Start)
})
l.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()
l.add(ban)
}
return added
}
// holds reports whether the ledger holds a ban on netblock that started
// at start.
func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
bans, found := l.netblocks.Peek(netblock)
return found && slices.ContainsFunc(*bans, func(ban Ban) bool {
return ban.Start.Equal(start)
})
}
// ban bans netblock at now for cause, with reason and notes, as
// BanForLimit and BanForAttack describe, and returns the ban, and whether
// it made it. Unless keep is true, the ban is not made, only returned: it
// is the ban that would have been made.
func (l *Ledger) ban(
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool,
) (Ban, bool) {
l.mu.Lock()
defer l.mu.Unlock()
// held are the netblock's bans, none of them active.
var held []Ban
bans, found := l.netblocks.Get(netblock)
if found {
active := activeBan(*bans, now)
if active != nil {
return *active, false
}
held = *bans
notes.EarlierBans = earlierBans(held)
}
notes.Request = notes.Request.cut()
ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes}
if cause == CauseAttack {
ban.Expires = l.attackExpiry(held, now)
} else {
ban.Expires = l.limitExpiry(held, now)
}
if !keep {
return ban, true
}
l.add(ban)
l.made[cause]++
l.markChanged()
return ban, true
}
// earlierBans returns how many bans a netblock with the bans held, oldest
// first, has had, by cause: the first ban held counts the bans the
// netblock had before that one, since dropped to make room, and each ban
// held adds one.
func earlierBans(held []Ban) EarlierBans {
earlier := held[0].Notes.EarlierBans
for _, ban := range held {
switch ban.Cause {
case CauseLimit:
earlier.Limit++
case CauseAttack:
earlier.Attack++
case CauseAdmin:
earlier.Admin++
}
}
return earlier
}
// markChanged has Changed receive a value, unless one is waiting already.
func (l *Ledger) markChanged() {
select {
case l.changed <- struct{}{}:
default: // a value is waiting already
}
}
// active returns the ban active at now on a netblock client is in, or
@@ -355,10 +632,32 @@ func (l *Ledger) active(client netip.Addr, now time.Time) *Ban {
return nil
}
// covering returns the bans of each netblock held that client is in,
// leaving when the netblocks were last seen unchanged.
func (l *Ledger) covering(client netip.Addr) []*[]Ban {
lengths := l.v6Lengths
if client.Is4() {
lengths = l.v4Lengths
}
var found []*[]Ban
for _, length := range lengths {
bans, ok := l.netblocks.Peek(netip.PrefixFrom(client, length).Masked())
if ok {
found = append(found, bans)
}
}
return found
}
// add adds ban to its netblock's bans, after the last, and makes its
// netblock the most recently seen. With MaxBans held, it drops one first.
// netblock the most recently seen. With MaxBans held, it drops one first,
// unless ban's cause is CauseAdmin, which does not count toward MaxBans.
func (l *Ledger) add(ban Ban) {
if l.held == l.rules.MaxBans {
counted := ban.Cause != CauseAdmin
if counted && l.held == l.rules.MaxBans {
l.dropOne()
}
@@ -371,7 +670,10 @@ func (l *Ledger) add(ban Ban) {
}
*bans = append(*bans, ban)
l.held++
if counted {
l.held++
}
lengths := &l.v6Lengths
if ban.Netblock.Addr().Is4() {
@@ -383,12 +685,24 @@ 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 ban that ended last, or nil
// when it has none.
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
// limitExpiry returns when a ban for a broken limit made at now ends, or
// zero when it is permanent. held are the netblock's bans, none of them
// active, of which the one that ended last, other than a ban for a clear
// sign of attack or a lifted one, can make the new ban longer. A ban an
// admin adds to bans.json can start after another and end before it, so
// that one is looked for among them all.
func (l *Ledger) limitExpiry(held []Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration
var last *Ban
for i, ban := range held {
if ban.Cause != CauseAttack && ban.Lifted.IsZero() &&
(last == nil || ban.Expires.After(last.Expires)) {
last = &held[i]
}
}
if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow {
lastLength := last.Expires.Sub(last.Start)
// This is repeatFactor * lastLength > MaxBanDuration, written so
@@ -407,17 +721,53 @@ func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
return now.Add(length)
}
// dropOne drops the earliest ban of the netblock that has gone longest
// without a request, and the netblock with it if that was its only ban.
func (l *Ledger) dropOne() {
netblock, bans, _ := l.netblocks.GetOldest()
if len(*bans) == 1 {
l.netblocks.Remove(netblock)
} else {
*bans = slices.Delete(*bans, 0, 1)
// attackExpiry returns when a ban for a clear sign of attack made at now
// ends. held are the netblock's bans, none of them active: if one of them
// is for a clear sign of attack too, and was not lifted, the new ban is
// permanent, and its end zero; otherwise it ends AttackBanDuration later.
func (l *Ledger) attackExpiry(held []Ban, now time.Time) time.Time {
for _, ban := range held {
if ban.Cause == CauseAttack && ban.Lifted.IsZero() {
return time.Time{}
}
}
l.held--
return now.Add(l.rules.AttackBanDuration)
}
// dropOne drops the earliest ban whose cause is not CauseAdmin of the
// netblock that has gone longest without a request, of those that hold
// such a ban, and the netblock with it if that was its only ban. It is
// called with at least one such ban held. It looks at each netblock once
// at most, and drops nothing when none holds such a ban.
func (l *Ledger) dropOne() {
for range l.netblocks.Len() {
netblock, bans, _ := l.netblocks.GetOldest()
i := slices.IndexFunc(*bans, func(ban Ban) bool {
return ban.Cause != CauseAdmin
})
if i < 0 {
// Its bans are all an admin's, and never dropped. Get makes
// it the most recently seen, so that the next netblock is
// looked at; when it was seen matters only for dropping a
// ban, and a ban added to it makes it the most recently seen
// anyway.
l.netblocks.Get(netblock)
continue
}
if len(*bans) == 1 {
l.netblocks.Remove(netblock)
} else {
*bans = slices.Delete(*bans, i, i+1)
}
l.held--
return
}
}
// cut returns r with each text cut to maxTextBytes and copied, so that
+189 -31
View File
@@ -21,11 +21,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
// 81 hours.
for i, hours := range []int{1, 3, 9, 27, 81} {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
length := time.Duration(hours) * time.Hour
if !ban.Expires.Equal(now.Add(length)) || ban.Notes.EarlierBans != i {
t.Fatalf("ban %d lasts %s with %d earlier bans, want %d hours and %d",
if !ban.Expires.Equal(now.Add(length)) ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: i}) {
t.Fatalf("ban %d lasts %s with earlier bans %+v, want %d hours and %d for a limit",
i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i)
}
@@ -34,12 +35,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// The sixth would last 243 hours, more than seven days: it is
// permanent, and never ends.
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
}
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day))
_, banned, _ := ledger.Check(netblock.Addr(), now.Add(100*365*day))
if !banned {
t.Error("a permanent ban ended")
}
@@ -63,11 +64,12 @@ func TestRepeatWindowRunsOut(t *testing.T) {
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second, _ := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
if second.Expires.Sub(second.Start) != tc.want || second.Notes.EarlierBans != 1 {
t.Errorf("second ban lasts %s with %d earlier bans, want %s and 1",
if second.Expires.Sub(second.Start) != tc.want ||
second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("second ban lasts %s with earlier bans %+v, want %s and 1 for a limit",
second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want)
}
})
@@ -81,7 +83,7 @@ func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
ledger := bans.New(rules)
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
ban, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
if !ban.Permanent() {
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
@@ -101,7 +103,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now := midnight()
for i := range 14 {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Expires.After(ban.Start) {
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
}
@@ -109,7 +111,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now = ban.Expires
}
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
}
@@ -121,12 +123,22 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
first, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
if !made {
t.Error("the first ban was not made")
}
if again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
again, len(ledger.Bans(netblock)), first)
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
}
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first {
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
again, made, first)
}
}
@@ -135,21 +147,21 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
for range 3 {
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
got, banned, _ := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got.Start != ban.Start {
t.Fatalf("check during the ban gives %+v and %t", got, banned)
}
}
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
_, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
if banned {
t.Error("another netblock is banned")
}
_, banned = ledger.Check(netblock.Addr(), ban.Expires)
_, banned, _ = ledger.Check(netblock.Addr(), ban.Expires)
if banned {
t.Error("the ban did not end")
}
@@ -167,14 +179,14 @@ func TestFindCountsNothing(t *testing.T) {
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
got, banned := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
got, banned, _ := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got != ban {
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
}
_, banned = ledger.Find(netblock.Addr(), ban.Expires)
_, banned, _ = ledger.Find(netblock.Addr(), ban.Expires)
if banned {
t.Error("the ban did not end")
}
@@ -196,7 +208,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
d := netip.MustParsePrefix("2001:db8::/64")
now := midnight()
first := ledger.BanForLimit(a, now, bans.Notes{})
first, _ := ledger.BanForLimit(a, now, bans.Notes{})
ledger.BanForLimit(b, now, bans.Notes{})
ledger.BanForLimit(c, now, bans.Notes{})
@@ -231,13 +243,158 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
held := ledger.Bans(netblock)
if len(held) != 1 || held[0] != second || held[0].Notes.EarlierBans != 1 {
t.Errorf("the ledger holds %+v, want only the second ban, with 1 earlier ban",
held)
if len(held) != 1 || held[0] != second ||
held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the ledger holds %+v, want only the second ban, "+
"with 1 earlier ban for a limit", held)
}
}
func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
notes := bans.Notes{RuleID: "env-file", Target: "path"}
ban, _ := ledger.BanForAttack(netblock, midnight(), notes)
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
ledger.Made(bans.CauseLimit) != 0 {
t.Fatalf("the ban is %+v, with %d made for an attack and %d for a limit, "+
"want one for an attack, of seven days", ban,
ledger.Made(bans.CauseAttack), ledger.Made(bans.CauseLimit))
}
wantChanged(t, ledger, true)
// In observe mode the ban refuses nothing, and stays as it is, while
// Find tells that the request would have made it permanent.
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent {
t.Fatalf("a request found under the ban left it %+v, would have made it "+
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
}
wantChanged(t, ledger, false)
// A request it refuses makes it permanent, says so, and makes
// bans.json due.
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
if !madePermanent || !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
t.Fatalf("after a request during the ban, it is %+v, made permanent %t, "+
"want it made permanent", got, madePermanent)
}
wantChanged(t, ledger, true)
// The next request finds it permanent already.
_, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
if !banned || madePermanent {
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
"want banned by the permanent ban", banned, madePermanent)
}
}
func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
// A ban for a broken limit before does not count.
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second, _ := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
if second.Expires.Sub(second.Start) != 7*day {
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
second.Expires.Sub(second.Start))
}
// Once that has run out without a request, the netblock is served, and
// its next clear sign of attack bans it for good.
_, banned, _ := ledger.Check(netblock.Addr(), second.Expires)
if banned {
t.Fatal("the ban did not end")
}
// Its notes show the earlier ban for an attack that makes it permanent,
// beside the one for a limit.
third, _ := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
if !third.Permanent() ||
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
"with 1 earlier ban for a limit and 1 for an attack", third)
}
}
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// While the first ban lasts, none would be made.
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
if would || during != first {
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
would, during, first)
}
// As it ends, a clear sign of attack would ban for seven days, and a
// limit broken again for three hours, but neither is made.
limitNotes := bans.Notes{Limit: 1, Window: "minute"}
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
bans.Notes{RuleID: "git-dir"})
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
if !wouldAttack || !attack.Expires.Equal(first.Expires.Add(7*day)) ||
attack.Reason != "matched the rule git-dir" || !wouldLimit ||
!limit.Expires.Equal(first.Expires.Add(3*time.Hour)) ||
limit.Reason != "requests per minute over the limit of 1" {
t.Errorf("would ban with %+v and %+v, want seven days for the attack and "+
"three hours for the limit", attack, limit)
}
if len(ledger.Bans(netblock)) != 1 || ledger.Made(bans.CauseLimit) != 1 ||
ledger.Made(bans.CauseAttack) != 0 {
t.Errorf("the ledger holds %+v, want the first ban alone", ledger.Bans(netblock))
}
wantChanged(t, ledger, false)
// The ban made is the one that would have been.
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
if made != limit {
t.Errorf("the ban made is %+v, want %+v", made, limit)
}
}
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
// Three times the seven days would be permanent; a limit broken as the
// ban for an attack ends bans for an hour, as a first broken limit does.
attack, _ := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
limit, _ := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
}
// And a request during the ban for a limit leaves it as it is.
got, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start)
if got.Permanent() || madePermanent {
t.Error("a request during a ban for a limit made it permanent")
}
}
@@ -251,7 +408,7 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
}
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
cut := long[:256]
want := bans.Request{
@@ -269,6 +426,7 @@ func defaultRules() bans.Rules {
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: day,
MaxBanDuration: 7 * day,
AttackBanDuration: 7 * day,
MaxBans: 5000,
}
}
+39 -22
View File
@@ -42,7 +42,7 @@ func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
high := netip.MustParsePrefix("203.0.113.10/32")
low := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(v6, midnight(), bans.Notes{})
first, _ := ledger.BanForLimit(v6, midnight(), bans.Notes{})
ledger.BanForLimit(high, midnight(), bans.Notes{})
ledger.BanForLimit(low, midnight(), bans.Notes{})
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
@@ -68,7 +68,7 @@ func TestLoadedBansCarryOn(t *testing.T) {
before := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
ban, _ := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
// Loaded into a new ledger, as across a restart, the ban still refuses
// while it lasts, and once it has ended a broken limit bans for three
@@ -76,14 +76,15 @@ func TestLoadedBansCarryOn(t *testing.T) {
after := bans.New(defaultRules())
after.Load(before.Snapshot())
_, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
_, banned, _ := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
if !banned {
t.Error("the loaded ban does not refuse")
}
again := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
if again.Expires.Sub(again.Start) != 3*time.Hour || again.Notes.EarlierBans != 1 {
t.Errorf("the next ban lasts %s with %d earlier bans, want 3h and 1",
again, _ := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
if again.Expires.Sub(again.Start) != 3*time.Hour ||
again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
again.Expires.Sub(again.Start), again.Notes.EarlierBans)
}
}
@@ -110,7 +111,7 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
"198.51.100.7": true,
"198.51.100.8": false,
} {
_, banned := ledger.Check(netip.MustParseAddr(client), midnight())
_, banned, _ := ledger.Check(netip.MustParseAddr(client), midnight())
if banned != want {
t.Errorf("%s is refused: %t, want %t", client, banned, want)
}
@@ -149,19 +150,19 @@ func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
now := midnight().Add(2 * time.Hour)
client := netip.MustParseAddr("203.0.113.9")
ban, banned := ledger.Find(client, now)
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)
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{})
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)))
@@ -172,13 +173,14 @@ 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.
// 1-hour ban added to bans.json over it, with no cause and no notes.
netblock := netip.MustParsePrefix("203.0.113.9/32")
nineHours := bans.Ban{
Netblock: netblock,
Start: midnight(),
Expires: midnight().Add(9 * time.Hour),
Notes: bans.Notes{EarlierBans: 2},
Cause: bans.CauseLimit,
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 2}},
}
admins := bans.Ban{
Netblock: netblock,
@@ -191,10 +193,12 @@ func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
// 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",
// before the 9-hour one and it, for a limit, and the admin's.
ban, _ := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
"want 27h, 3 for a limit and 1 an admin's",
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
}
}
@@ -203,10 +207,15 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
t.Parallel()
// bans.json lists the bans by netblock, not in the order they began.
later := bans.Ban{Netblock: netip.MustParsePrefix("203.0.113.1/32"), Start: midnight()}
later := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.1/32"),
Start: midnight(),
Cause: bans.CauseLimit,
}
earlier := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
Start: midnight().Add(-time.Hour),
Cause: bans.CauseLimit,
}
rules := defaultRules()
@@ -228,9 +237,17 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
rules := defaultRules()
rules.MaxBans = 3
ledger := bans.New(rules)
kept := bans.Ban{Netblock: netip.MustParsePrefix("2001:db8::/64"), Start: midnight()}
kept := bans.Ban{
Netblock: netip.MustParsePrefix("2001:db8::/64"),
Start: midnight(),
Cause: bans.CauseLimit,
}
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.0/24"), Start: midnight()},
{
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Start: midnight(),
Cause: bans.CauseLimit,
},
kept,
})
@@ -238,15 +255,15 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
// bans.json is taken in, that ban is lifted.
ledger.Load([]bans.Ban{kept})
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
_, 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(),
first, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
bans.Notes{})
second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
second, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
bans.Notes{})
want := []bans.Ban{first, second, kept}
+547 -23
View File
@@ -1,9 +1,11 @@
// Package config reads smallwebwaf's settings. Every setting is an
// environment variable whose name starts with SWWAF_, every setting has a
// default, and this package is the one place they are read.
// environment variable whose name starts with SWWAF_, or a file such a
// variable names, every setting has a default, and this package is the
// one place they are read.
package config
import (
"crypto/x509"
"errors"
"fmt"
"log/slog"
@@ -12,12 +14,16 @@ import (
"net/http"
"net/netip"
"net/url"
"os"
"path/filepath"
"slices"
"strconv"
"strings"
"time"
"unicode/utf8"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
@@ -27,10 +33,14 @@ type Config struct {
ListenAddr string
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
UpstreamURL *url.URL
// InstanceName is the name each request log line gives as instance
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets
// to the first 12 characters of the container's id.
InstanceName string
// Observe is true in observe mode, when SWWAF_MODE is observe rather
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country
// lists or a rate limit would refuse is passed to the app instead, and
// no ban is made.
// lists, a rate limit or a rule would refuse is passed to the app
// instead, and no ban is made.
Observe bool
// TrustedProxies are the netblocks whose X-Forwarded-For is
// believed (SWWAF_TRUSTED_PROXIES).
@@ -75,6 +85,10 @@ type Config struct {
RateLimitPerMinute int64
RateLimitPerHour int64
RateLimitPerDay int64
// RateLimitExemptPaths are the path prefixes whose requests the rate
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
// Each starts with /.
RateLimitExemptPaths []string
// DeniedCountries are the countries whose clients are refused
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
// empty, are the only countries whose clients are let through
@@ -85,7 +99,8 @@ type Config struct {
// BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
// breaks a rate limit, SWWAF_DENY_NETS and the country lists.
// breaks a rate limit or matches a ban rule, SWWAF_DENY_NETS and the
// country lists.
BanResponse int
// LimitBanDuration is the ban for a first broken rate limit
// (SWWAF_LIMIT_BAN_DURATION). A limit broken again within
@@ -96,6 +111,9 @@ type Config struct {
LimitBanDuration time.Duration
LimitBanRepeatWindow time.Duration
MaxBanDuration time.Duration
// AttackBanDuration is the ban for a first clear sign of attack
// (SWWAF_ATTACK_BAN_DURATION). It cannot be off.
AttackBanDuration time.Duration
// MaxBans is the most bans held (SWWAF_MAX_BANS).
MaxBans int
// BanScopeV4Prefix is the length of the netblock around an IPv4
@@ -109,15 +127,54 @@ type Config struct {
StateDir string
StateWriteDelay time.Duration
StateCounterInterval time.Duration
// LogRequestHeaders are the request headers whose values the request
// log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS).
LogRequestHeaders []string
// AdminToken is the bearer token an admin sends for the ban endpoints
// and /_smallwebwaf/clients/<ip> (SWWAF_ADMIN_TOKEN), "" while it is
// unset and they are off.
AdminToken string
// MetricsToken is the bearer token a scraper sends for the metrics
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
// MetricsTopN is how many countries get series of their own in the
// metrics (SWWAF_METRICS_TOP_N).
MetricsToken string
MetricsTopN int
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
// unless RulesEnabled is false (SWWAF_RULES_ENABLED).
RulesDir string
RulesEnabled bool
// LogRemoteURL is where every line on stdout is also sent
// (SWWAF_LOG_REMOTE_URL), nil while it is unset and nothing is sent.
// LogRemoteTLSCAs are the certificates a syslog+tls endpoint's
// certificate must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil while
// it is unset and the host's own are used. LogRemoteBuffer is the most
// lines held while they wait to be sent (SWWAF_LOG_REMOTE_BUFFER).
// LogRemoteFacility is the number of the syslog facility
// (SWWAF_LOG_REMOTE_FACILITY), and LogRemoteAppName the APP-NAME
// (SWWAF_LOG_REMOTE_APP_NAME, by default InstanceName), of the records
// the lines are sent in.
LogRemoteURL *url.URL
LogRemoteTLSCAs *x509.CertPool
LogRemoteBuffer int
LogRemoteFacility int
LogRemoteAppName string
// AlertWebhookURL is where each alert is posted as JSON
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset and no alert is
// sent. AlertWebhookHeaders are sent with each
// (SWWAF_ALERT_WEBHOOK_HEADERS). AlertEvents are the events alerts are
// sent for (SWWAF_ALERT_EVENTS). A repeat of an alert within
// AlertCooldown is held back (SWWAF_ALERT_COOLDOWN), and so is an alert
// past AlertMaxPerHour in an hour, for the hour's summary
// (SWWAF_ALERT_MAX_PER_HOUR); 0 is off for both.
AlertWebhookURL *url.URL
AlertWebhookHeaders http.Header
AlertEvents []string
AlertCooldown time.Duration
AlertMaxPerHour int
// settings are the values read, as given or by default, for the
// log line at start.
// settings are the values read, as given or by default, and the
// files they were read from, for the log line at start.
settings []slog.Attr
}
@@ -133,8 +190,13 @@ const (
ipv4Bits = 32
// minTokenLength is the fewest characters a token may have.
minTokenLength = 32
// masked is what the log shows for a token that is set.
// masked is what the log shows for a token that is set, and in place of
// a secret in another setting.
masked = "********"
// defaultListenAddr and defaultUpstreamURL are the defaults of
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
defaultListenAddr = ":8080"
defaultUpstreamURL = "http://127.0.0.1:8081"
)
var (
@@ -155,6 +217,11 @@ var (
"such as http://127.0.0.1:8081")
errNotCountry = errors.New(
"is not a two-letter country code such as de or kp")
errNotHeaderName = errors.New(
"is not a header name such as accept-language")
errHeaderTakenOut = errors.New(
"is taken out of every request by Go's HTTP server, so it can never " +
"be logged")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New(
@@ -166,18 +233,43 @@ 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")
errNotPathPrefix = errors.New(
"is not a path prefix starting with /, such as /assets/")
errNotBoolean = errors.New("is not true or false")
errNotLogRemoteURL = errors.New(
"is not syslog+udp, syslog+tcp or syslog+tls with a host and a port, " +
"and nothing more, such as syslog+tls://logs.example:6514")
errNoCertificate = errors.New("holds no PEM certificate")
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
errNotAppName = errors.New(
"is not 1 to 48 printable ASCII characters without a space, such as gitea")
errSetTwice = errors.New("set only one of them")
errNotWebhookURL = errors.New(
"is not an http or https URL without a user or a fragment, " +
"such as https://alerts.example/smallwebwaf")
errNotWebhookHeader = errors.New(
"is not a header name followed by : and the header's value, " +
"such as Authorization:Bearer <token>")
errNotAlertEvent = errors.New(
"is not ban, permanent_ban, waf_block, anomaly, reputation_hit, " +
"source_failure or file_error")
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
)
// FromEnvironment reads the settings with lookupEnv, normally
// os.LookupEnv. A setting that is not set takes its default. A setting
// that is set but invalid is an error that names it.
// os.LookupEnv. A setting may instead be given as a file: the variable
// named by the setting's name with _FILE added names the file, which is
// read now (see lookup). A setting that is not set takes its default. A
// setting that is set but invalid is an error that names it.
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname),
Observe: env.observe("SWWAF_MODE", "enforce"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
@@ -195,6 +287,7 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"),
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
@@ -202,15 +295,34 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
AttackBanDuration: env.durationNotOff("SWWAF_ATTACK_BAN_DURATION", "7d"),
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
"accept,accept-language,accept-encoding,content-type,origin,range"),
AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
AlertWebhookURL: env.webhookURL("SWWAF_ALERT_WEBHOOK_URL"),
AlertWebhookHeaders: env.webhookHeaders("SWWAF_ALERT_WEBHOOK_HEADERS"),
AlertEvents: env.alertEvents("SWWAF_ALERT_EVENTS",
strings.Join(alerts.Events(), ",")),
AlertCooldown: env.duration("SWWAF_ALERT_COOLDOWN", "15m"),
AlertMaxPerHour: env.numberOrOff("SWWAF_ALERT_MAX_PER_HOUR", "60"),
}
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
cfg.InstanceName, cfg.LogRemoteURL != nil)
for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) {
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
@@ -227,6 +339,24 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
return cfg, nil
}
// ListenAddrAndUpstreamURL reads only SWWAF_LISTEN_ADDR and
// SWWAF_UPSTREAM_URL, either of which may be given as a file, as
// FromEnvironment does. The health check needs no other setting, so it
// reads no other, nor a file that another names.
func ListenAddrAndUpstreamURL(
lookupEnv func(string) (string, bool),
) (string, *url.URL, error) {
env := &environment{lookupEnv: lookupEnv}
listenAddr := env.address("SWWAF_LISTEN_ADDR", defaultListenAddr)
upstreamURL := env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL)
if env.err != nil {
return "", nil, env.err
}
return listenAddr, upstreamURL, nil
}
// privateRanges are the private address ranges, the default trusted
// proxies.
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
@@ -248,8 +378,8 @@ type environment struct {
// value returns a setting's value, or its default when it is not set,
// and notes it for the log.
func (e *environment) value(name, defaultValue string) string {
value, ok := e.lookupEnv(name)
if !ok {
value, set := e.lookup(name)
if !set {
value = defaultValue
}
@@ -258,6 +388,37 @@ func (e *environment) value(name, defaultValue string) string {
return value
}
// lookup returns a setting's value and whether it is set: the value of the
// variable name, or the contents of the file that the variable name_FILE
// names, less one newline at their end. It notes that file's path for the
// log. Both variables set, or a file that cannot be read, is an error.
func (e *environment) lookup(name string) (string, bool) {
value, set := e.lookupEnv(name)
fileName := name + "_FILE"
path, inFile := e.lookupEnv(fileName)
if !inFile {
return value, set
}
if set {
e.check(name, fmt.Errorf("is set, and so is %s; %w", fileName, errSetTwice))
return value, set
}
e.settings = append(e.settings, slog.String(fileName, path))
contents, err := os.ReadFile(path) //nolint:gosec // a file the admin names
if err != nil {
e.check(fileName, fmt.Errorf("cannot be read: %w", err))
return "", false
}
return strings.TrimSuffix(string(contents), "\n"), true
}
// check keeps the first error, naming the setting it is about.
func (e *environment) check(name string, err error) {
if err != nil && e.err == nil {
@@ -292,6 +453,16 @@ func (e *environment) observe(name, defaultValue string) bool {
return mode == "observe"
}
// boolean reads a setting that is true or false.
func (e *environment) boolean(name, defaultValue string) bool {
value := e.value(name, defaultValue)
if value != "true" && value != "false" {
e.check(name, fmt.Errorf("%q %w", value, errNotBoolean))
}
return value == "true"
}
// netblocks reads a setting that is a list of netblocks.
func (e *environment) netblocks(name, defaultValue string) []netip.Prefix {
netblocks, err := parseNetblocks(e.value(name, defaultValue))
@@ -333,6 +504,14 @@ func (e *environment) count(name, defaultValue string) int64 {
return count
}
// pathPrefixes reads a setting that is a list of path prefixes.
func (e *environment) pathPrefixes(name, defaultValue string) []string {
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
e.check(name, err)
return prefixes
}
// countries reads a setting that is a list of countries.
func (e *environment) countries(name, defaultValue string) []string {
countries, err := parseCountries(e.value(name, defaultValue))
@@ -341,10 +520,19 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries
}
// headerNames reads a setting that is a list of header names, and
// returns them in lower case.
func (e *environment) headerNames(name, defaultValue string) []string {
headers, err := parseHeaderNames(e.value(name, defaultValue))
e.check(name, err)
return headers
}
// durationNotOff reads a setting that is a duration and, unlike a
// timeout, cannot be off.
func (e *environment) durationNotOff(name, defaultValue string) time.Duration {
duration, err := parseDurationNotOff(e.value(name, defaultValue))
duration, err := ParseDurationNotOff(e.value(name, defaultValue))
e.check(name, err)
return duration
@@ -389,7 +577,7 @@ func (e *environment) absolutePath(name, defaultValue string) string {
// switches off what it guards; set, it must be at least minTokenLength
// characters. Neither the log nor an error shows its value.
func (e *environment) token(name string) string {
value, set := e.lookupEnv(name)
value, set := e.lookup(name)
if !set {
e.settings = append(e.settings, slog.String(name, ""))
@@ -405,6 +593,122 @@ 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. Its value names a file
// already, so, unlike the other settings, it has no _FILE form.
func (e *environment) certificates(name string) *x509.CertPool {
path, _ := e.lookupEnv(name)
e.settings = append(e.settings, slog.String(name, path))
if path == "" {
return nil
}
pem, err := os.ReadFile(path) //nolint:gosec // a file the admin names
if err != nil {
e.check(name, fmt.Errorf("cannot be read: %w", err))
return nil
}
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(pem) {
e.check(name, fmt.Errorf("%q %w", path, errNoCertificate))
return nil
}
return pool
}
// facility reads a setting that is a syslog facility, and returns its
// number.
func (e *environment) facility(name, defaultValue string) int {
number, err := parseFacility(e.value(name, defaultValue))
e.check(name, err)
return number
}
// appName reads the setting that is the APP-NAME of the records the log
// lines are sent in, by default the instance name. Its value is checked
// when it is set, and, while lines are sent, when it is the instance name.
func (e *environment) appName(name, instanceName string, sending bool) string {
value, set := e.lookup(name)
if !set {
value = instanceName
}
e.settings = append(e.settings, slog.String(name, value))
switch {
case isAppName(value):
case set:
e.check(name, fmt.Errorf("%q %w", value, errNotAppName))
case sending:
e.check(name, fmt.Errorf("is unset, and SWWAF_INSTANCE_NAME %q, its default, %w",
value, errNotAppName))
}
return value
}
// webhookURL reads the setting that is where each alert is posted. Unset
// or empty, it is nil, and no alert is sent. The log shows ******** in
// place of its path and query, and an error shows none of it, since many
// webhooks carry their secret there.
func (e *environment) webhookURL(name string) *url.URL {
value, _ := e.lookup(name)
webhook, logged, err := parseWebhookURL(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return webhook
}
// webhookHeaders reads the setting that is the headers sent with each
// alert. The log shows each header's value as ********, since a header
// such as Authorization carries a secret.
func (e *environment) webhookHeaders(name string) http.Header {
value, _ := e.lookup(name)
headers, logged, err := parseWebhookHeaders(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return headers
}
// alertEvents reads the setting that is the events alerts are sent for.
func (e *environment) alertEvents(name, defaultValue string) []string {
events, err := parseAlertEvents(e.value(name, defaultValue))
e.check(name, err)
return events
}
// numberOrOff reads a setting that is a whole number above zero, or off,
// which is 0.
func (e *environment) numberOrOff(name, defaultValue string) int {
number, err := parseNumberOrOff(e.value(name, defaultValue))
e.check(name, err)
return number
}
// 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) {
@@ -506,9 +810,9 @@ func parseCount(value string) (int64, error) {
return n, nil
}
// parseDurationNotOff reads a duration above zero, as parseDuration does,
// but not off.
func parseDurationNotOff(value string) (time.Duration, error) {
// ParseDurationNotOff reads a duration above zero, as parseDuration does,
// but not off. The ban endpoint reads the duration of a ban with it too.
func ParseDurationNotOff(value string) (time.Duration, error) {
duration, err := parseDuration(value)
if err != nil || duration == 0 {
return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero)
@@ -611,6 +915,23 @@ func parseNetblock(value string) (netip.Prefix, error) {
return netip.PrefixFrom(addr, addr.BitLen()), nil
}
// parsePathPrefixes reads a comma-separated list of path prefixes, each
// starting with /.
func parsePathPrefixes(value string) ([]string, error) {
prefixes, err := parseList(value)
if err != nil {
return nil, err
}
for _, prefix := range prefixes {
if !strings.HasPrefix(prefix, "/") {
return nil, fmt.Errorf("%q %w", prefix, errNotPathPrefix)
}
}
return prefixes, nil
}
// countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK,
// the code in common use for Kosovo. golang.org/x/text/language cannot
// check them: it also takes withdrawn codes such as su, and reserved ones
@@ -667,6 +988,58 @@ func parseCountries(value string) ([]string, error) {
return countries, nil
}
// headerNameChars are the characters RFC 9110 allows in a header name:
// letters, digits and these marks.
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
"0123456789!#$%&'*+-.^_`|~"
// IsHeaderName reports whether name can be a header name: one or more of
// the characters RFC 9110 allows in one.
func IsHeaderName(name string) bool {
if name == "" {
return false
}
for _, char := range name {
if !strings.ContainsRune(headerNameChars, char) {
return false
}
}
return true
}
// parseHeaderNames reads a comma-separated list of header names in either
// case, and returns them in lower case. Host and Transfer-Encoding are
// refused: Go's HTTP server takes them out of the request's headers.
func parseHeaderNames(value string) ([]string, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
headers := make([]string, 0, len(items))
for _, item := range items {
if !IsHeaderName(item) {
return nil, fmt.Errorf("%q %w", item, errNotHeaderName)
}
header := strings.ToLower(item)
switch header {
case "host":
return nil, fmt.Errorf("%q %w; the request's host is the field host",
item, errHeaderTakenOut)
case "transfer-encoding":
return nil, fmt.Errorf("%q %w", item, errHeaderTakenOut)
}
headers = append(headers, header)
}
return headers, nil
}
// parseListenAddr checks an address to listen on: an optional host and a
// port number.
func parseListenAddr(value string) (string, error) {
@@ -709,3 +1082,154 @@ 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
}
// parseWebhookURL reads where each alert is posted: http or https, a
// host, and an optional port from 1 to 65535, path and query, without a
// user or a fragment. It returns the URL, and how the log shows it: its
// scheme and host, and ******** in place of its path and query, if it has
// either. An error shows no part of the value. An empty value is no URL.
func parseWebhookURL(value string) (*url.URL, string, error) {
if value == "" {
return nil, "", nil
}
webhook, err := url.Parse(value)
if err != nil {
return nil, "", errNotWebhookURL
}
port, err := strconv.ParseUint(webhook.Port(), 10, 16)
valid := (webhook.Scheme == "http" || webhook.Scheme == "https") &&
webhook.Hostname() != "" && (webhook.Port() == "" || (err == nil && port != 0)) &&
webhook.User == nil && webhook.Opaque == "" && webhook.Fragment == ""
if !valid {
return nil, "", errNotWebhookURL
}
logged := webhook.Scheme + "://" + webhook.Host
if webhook.Path != "" || webhook.RawQuery != "" {
logged += "/" + masked
}
return webhook, logged, nil
}
// parseWebhookHeaders reads a comma-separated list of headers, each its
// name, :, and its value, and returns them, and how the log shows them,
// with each value as ********. An error names the item by its place in
// the list, so that it shows no value. An empty value is an empty list.
func parseWebhookHeaders(value string) (http.Header, string, error) {
headers := http.Header{}
if strings.TrimSpace(value) == "" {
return headers, "", nil
}
logged := []string{}
for i, item := range strings.Split(value, ",") {
name, headerValue, found := strings.Cut(item, ":")
name = strings.TrimSpace(name)
if !found || !IsHeaderName(name) || strings.ContainsAny(headerValue, "\r\n\x00") {
return nil, "", fmt.Errorf("item %d %w", i+1, errNotWebhookHeader)
}
headers.Add(name, strings.TrimSpace(headerValue))
logged = append(logged, name+":"+masked)
}
return headers, strings.Join(logged, ","), nil
}
// parseAlertEvents reads a comma-separated list of the events alerts can
// be sent for.
func parseAlertEvents(value string) ([]string, error) {
events, err := parseList(value)
if err != nil {
return nil, err
}
for _, event := range events {
if !slices.Contains(alerts.Events(), event) {
return nil, fmt.Errorf("%q %w", event, errNotAlertEvent)
}
}
return events, nil
}
// parseNumberOrOff reads a whole number above zero, or off, which is 0.
func parseNumberOrOff(value string) (int, error) {
if value == off {
return 0, nil
}
n, err := strconv.Atoi(value)
if err != nil || n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotNumberOrOff)
}
return n, 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
}
+664 -28
View File
@@ -2,10 +2,15 @@ package config_test
import (
"bytes"
"crypto/x509"
"encoding/json"
"log/slog"
"maps"
"net/http"
"net/netip"
"os"
"path/filepath"
"reflect"
"slices"
"strings"
"testing"
@@ -34,23 +39,77 @@ const (
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
rulesDir = "SWWAF_RULES_DIR"
rulesEnabled = "SWWAF_RULES_ENABLED"
logRemoteURL = "SWWAF_LOG_REMOTE_URL"
logRemoteTLSCAFile = "SWWAF_LOG_REMOTE_TLS_CA_FILE"
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertWebhookHeaders = "SWWAF_ALERT_WEBHOOK_HEADERS"
alertEvents = "SWWAF_ALERT_EVENTS"
alertCooldown = "SWWAF_ALERT_COOLDOWN"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
)
// token is a token of 32 characters, the shortest allowed.
const token = "0123456789abcdef0123456789abcdef"
// defaultAlertEvents is the default of SWWAF_ALERT_EVENTS, and
// defaultAlertCooldown that of SWWAF_ALERT_COOLDOWN.
const (
defaultAlertEvents = "ban,permanent_ban,waf_block,anomaly,reputation_hit," +
"source_failure,file_error"
defaultAlertCooldown = "15m"
)
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," +
"content-type,origin,range"
// testCA is a CA certificate, of which only that it reads matters here.
const testCA = `-----BEGIN CERTIFICATE-----
MIIBkzCCATmgAwIBAgIUeySaE27dnr6A2HijrMB13gTLUKIwCgYIKoZIzj0EAwIw
HjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBDQTAgFw0yNjEwMDYxNDI2MTNa
GA8yMTI2MDkxMjE0MjYxM1owHjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBD
QTBZMBMGByqGSM49AgEGCCqGSM49AwEHA0IABKDEhcWKKhet2KgSdME+iEPxyEyn
2sd9IdElbt8DM2SfCdB2JsXo0C07UNZaywMPMfn/n8LNI/PKwu+N2uX7gfSjUzBR
MB0GA1UdDgQWBBRZo3BPLv0KbV4drw6JI1JIUKmoRzAfBgNVHSMEGDAWgBRZo3BP
Lv0KbV4drw6JI1JIUKmoRzAPBgNVHRMBAf8EBTADAQH/MAoGCCqGSM49BAMCA0gA
MEUCICC5k+76UpWoSwVbZA+atu5WcALEOJGqwOUWua3zemhcAiEA8Hfdxgwp0z2v
rlG9y/jrJb6ORy3kTLWo2EA0BA67vuI=
-----END CERTIFICATE-----
`
// token is a token of 32 characters, the shortest allowed, and
// otherToken another.
const (
token = "0123456789abcdef0123456789abcdef"
otherToken = "fedcba9876543210fedcba9876543210"
)
// instance is an SWWAF_INSTANCE_NAME that is a valid app name too, and
// remoteURL an SWWAF_LOG_REMOTE_URL, for the tests that send the lines.
const (
instance = "fsn1app1/gitea"
remoteURL = "syslog+udp://192.0.2.1:514"
)
// off switches a timeout, a size limit or a rate limit off.
const off = "off"
@@ -100,6 +159,7 @@ func TestDefaults(t *testing.T) {
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
AttackBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
StateDir: "/var/lib/smallwebwaf",
@@ -107,6 +167,8 @@ func TestDefaults(t *testing.T) {
StateCounterInterval: 15 * time.Minute,
MetricsToken: "",
MetricsTopN: 50,
RulesDir: "/etc/smallwebwaf/rules.d",
RulesEnabled: true,
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -120,6 +182,22 @@ func TestDefaults(t *testing.T) {
wantNetblocks(t, cfg.DenyNets)
wantCountries(t, deniedCountries, cfg.DeniedCountries)
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries)
hostname, err := os.Hostname()
if err != nil || hostname == "" || cfg.InstanceName != hostname {
t.Errorf("%s is %q, want the host's name %q (%v)", instanceName,
cfg.InstanceName, hostname, err)
}
wantHeaders := strings.Split(defaultLogRequestHeaders, ",")
if !slices.Equal(cfg.LogRequestHeaders, wantHeaders) {
t.Errorf("%s gave %v, want %v", logRequestHeaders, cfg.LogRequestHeaders,
wantHeaders)
}
if len(cfg.RateLimitExemptPaths) != 0 {
t.Errorf("%s gave %v, want none", rateLimitExemptPaths, cfg.RateLimitExemptPaths)
}
}
func TestValuesAsSet(t *testing.T) {
@@ -150,6 +228,7 @@ func TestValuesAsSet(t *testing.T) {
limitBanDuration: "15m",
limitBanRepeatWindow: "2d",
maxBanDuration: "30d",
attackBanDuration: "1d",
maxBans: "100",
banScopeV4Prefix: "24",
stateDir: "/srv/waf-state",
@@ -157,6 +236,8 @@ func TestValuesAsSet(t *testing.T) {
stateCounterInterval: "1h",
metricsToken: token,
metricsTopN: "10",
rulesDir: "/srv/waf-rules",
rulesEnabled: "false",
})
wantSettings(t, cfg, config.Config{
@@ -177,6 +258,7 @@ func TestValuesAsSet(t *testing.T) {
LimitBanDuration: 15 * time.Minute,
LimitBanRepeatWindow: 48 * time.Hour,
MaxBanDuration: 30 * 24 * time.Hour,
AttackBanDuration: 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
StateDir: "/srv/waf-state",
@@ -184,6 +266,8 @@ func TestValuesAsSet(t *testing.T) {
StateCounterInterval: time.Hour,
MetricsToken: token,
MetricsTopN: 10,
RulesDir: "/srv/waf-rules",
RulesEnabled: false,
})
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -198,6 +282,360 @@ func TestValuesAsSet(t *testing.T) {
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
}
func TestRateLimitExemptPathsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{rateLimitExemptPaths: "/assets/, /favicon.ico"})
if !slices.Equal(cfg.RateLimitExemptPaths, []string{"/assets/", "/favicon.ico"}) {
t.Errorf("%s gave %v, want /assets/ and /favicon.ico",
rateLimitExemptPaths, cfg.RateLimitExemptPaths)
}
}
func TestPathPrefixNotStartingWithSlashStopsTheStart(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(
environment{rateLimitExemptPaths: "/favicon.ico,assets/"}.lookupEnv)
want := rateLimitExemptPaths + `: "assets/" is not a path prefix ` +
`starting with /, such as /assets/`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
instanceName: "fsn1app1/gitea",
logRequestHeaders: " Accept , X-Custom",
})
if cfg.InstanceName != "fsn1app1/gitea" ||
!slices.Equal(cfg.LogRequestHeaders, []string{"accept", "x-custom"}) {
t.Errorf("%s is %q and %s gives %v", instanceName, cfg.InstanceName,
logRequestHeaders, cfg.LogRequestHeaders)
}
}
func TestRemoteLogSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{instanceName: instance})
if cfg.LogRemoteURL != nil || cfg.LogRemoteTLSCAs != nil ||
cfg.LogRemoteBuffer != 10000 || cfg.LogRemoteFacility != 16 ||
cfg.LogRemoteAppName != instance {
t.Errorf("remote log settings %v, %v, %d, %d and %q, want no URL, no "+
"certificates, 10000, 16 and %s's %s", cfg.LogRemoteURL,
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
cfg.LogRemoteAppName, instanceName, instance)
}
}
func TestRemoteLogSettingsAsSet(t *testing.T) {
t.Parallel()
caFile := filepath.Join(t.TempDir(), "ca.pem")
err := os.WriteFile(caFile, []byte(testCA), 0o600)
if err != nil {
t.Fatalf("write %s: %v", caFile, err)
}
cfg := fromEnvironment(t, environment{
logRemoteURL: "syslog+tls://logs.example:6514",
logRemoteTLSCAFile: caFile,
logRemoteBuffer: "500",
logRemoteFacility: "daemon",
logRemoteAppName: instance,
})
roots := x509.NewCertPool()
roots.AppendCertsFromPEM([]byte(testCA))
if cfg.LogRemoteURL.String() != "syslog+tls://logs.example:6514" ||
!roots.Equal(cfg.LogRemoteTLSCAs) || cfg.LogRemoteBuffer != 500 ||
cfg.LogRemoteFacility != 3 || cfg.LogRemoteAppName != instance {
t.Errorf("remote log settings %v, %v, %d, %d and %q", cfg.LogRemoteURL,
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
cfg.LogRemoteAppName)
}
}
func TestRemoteLogURLForms(t *testing.T) {
t.Parallel()
for _, value := range []string{
"syslog+udp://192.0.2.1:514",
"syslog+tcp://[2001:db8::1]:514",
"syslog+tls://logs.example:6514/",
} {
cfg := fromEnvironment(t, environment{logRemoteURL: value})
if cfg.LogRemoteURL.String() != value {
t.Errorf("%s read as %v", value, cfg.LogRemoteURL)
}
}
cfg := fromEnvironment(t, environment{logRemoteURL: ""})
if cfg.LogRemoteURL != nil {
t.Errorf("set but empty, %s read as %v", logRemoteURL, cfg.LogRemoteURL)
}
}
func TestRemoteLogFacilitiesByNumber(t *testing.T) {
t.Parallel()
for name, number := range map[string]int{
"kern": 0, "user": 1, "auth": 4, "authpriv": 10, "ftp": 11,
"local0": 16, "local5": 21, "local7": 23,
} {
cfg := fromEnvironment(t, environment{logRemoteFacility: name})
if cfg.LogRemoteFacility != number {
t.Errorf("%s read as %d, want %d", name, cfg.LogRemoteFacility, number)
}
}
}
func TestInvalidRemoteLogSettingStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct{ name, value string }{
{logRemoteURL, "logs.example:514"},
{logRemoteURL, "syslog://logs.example:514"},
{logRemoteURL, "http://logs.example:514"},
{logRemoteURL, "syslog+udp://logs.example"},
{logRemoteURL, "syslog+tcp://:514"},
{logRemoteURL, "syslog+tcp://logs.example:0"},
{logRemoteURL, "syslog+tls://logs.example:65536"},
{logRemoteURL, "syslog+tls://user@logs.example:6514"},
{logRemoteURL, "syslog+tcp://logs.example:514/app"},
{logRemoteURL, "syslog+tcp://logs.example:514?tls=1"},
{logRemoteTLSCAFile, "/nonexistent/ca.pem"},
{logRemoteBuffer, off}, {logRemoteBuffer, "0"}, {logRemoteBuffer, "10K"},
{logRemoteFacility, "local8"}, {logRemoteFacility, "LOCAL0"},
{logRemoteFacility, "16"}, {logRemoteFacility, ""},
{logRemoteAppName, ""}, {logRemoteAppName, "my app"},
{logRemoteAppName, "gitéa"}, {logRemoteAppName, strings.Repeat("a", 49)},
} {
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
if err == nil || !strings.HasPrefix(err.Error(), tc.name+": ") {
t.Errorf("%s=%q: error %v, want one naming it", tc.name, tc.value, err)
}
}
}
func TestRemoteLogCAFileWithoutCertificateStopsTheStart(t *testing.T) {
t.Parallel()
caFile := filepath.Join(t.TempDir(), "ca.pem")
err := os.WriteFile(caFile, []byte("not a certificate\n"), 0o600)
if err != nil {
t.Fatalf("write %s: %v", caFile, err)
}
_, err = config.FromEnvironment(environment{logRemoteTLSCAFile: caFile}.lookupEnv)
want := logRemoteTLSCAFile + `: "` + caFile + `" holds no PEM certificate`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestInstanceNameNotAnAppNameStopsTheStartOnlyWhileSending(t *testing.T) {
t.Parallel()
const spaced = "fsn1 app1"
sending := environment{logRemoteURL: remoteURL, instanceName: spaced}
_, err := config.FromEnvironment(sending.lookupEnv)
want := logRemoteAppName + `: is unset, and ` + instanceName +
` "fsn1 app1", its default, is not 1 to 48 printable ASCII characters ` +
`without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
cfg := fromEnvironment(t, environment{instanceName: spaced})
if cfg.LogRemoteAppName != spaced {
t.Errorf("not sending, %s is %q", logRemoteAppName, cfg.LogRemoteAppName)
}
sending[logRemoteAppName] = instance
cfg = fromEnvironment(t, sending)
if cfg.LogRemoteAppName != instance {
t.Errorf("set to %s, %s is %q", instance, logRemoteAppName,
cfg.LogRemoteAppName)
}
}
func TestAppNameSetStopsTheStartWhileSending(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
logRemoteURL: remoteURL,
instanceName: instance,
logRemoteAppName: "my app",
}.lookupEnv)
want := logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
`characters without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestAlertSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
if cfg.AlertWebhookURL != nil || len(cfg.AlertWebhookHeaders) != 0 ||
strings.Join(cfg.AlertEvents, ",") != defaultAlertEvents ||
cfg.AlertCooldown != 15*time.Minute || cfg.AlertMaxPerHour != 60 {
t.Errorf("alert settings %v, %v, %v, %s and %d, want no URL, no headers, "+
"%s, 15m and 60", cfg.AlertWebhookURL, cfg.AlertWebhookHeaders,
cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour, defaultAlertEvents)
}
}
func TestAlertSettingsAsSet(t *testing.T) {
t.Parallel()
const webhook = "https://alerts.example:8443/hooks/waf?team=ops"
cfg := fromEnvironment(t, environment{
alertWebhookURL: webhook,
alertWebhookHeaders: "Authorization: Bearer abc:def , x-team:ops",
alertEvents: "ban, file_error",
alertCooldown: "1h",
alertMaxPerHour: "10",
})
headers := http.Header{"Authorization": {"Bearer abc:def"}, "X-Team": {"ops"}}
if cfg.AlertWebhookURL.String() != webhook ||
!reflect.DeepEqual(cfg.AlertWebhookHeaders, headers) ||
!slices.Equal(cfg.AlertEvents, []string{"ban", "file_error"}) ||
cfg.AlertCooldown != time.Hour || cfg.AlertMaxPerHour != 10 {
t.Errorf("alert settings %v, %v, %v, %s and %d", cfg.AlertWebhookURL,
cfg.AlertWebhookHeaders, cfg.AlertEvents, cfg.AlertCooldown,
cfg.AlertMaxPerHour)
}
cfg = fromEnvironment(t, environment{
alertWebhookURL: "", alertEvents: "", alertCooldown: off, alertMaxPerHour: off,
})
if cfg.AlertWebhookURL != nil || len(cfg.AlertEvents) != 0 ||
cfg.AlertCooldown != 0 || cfg.AlertMaxPerHour != 0 {
t.Errorf("set empty or off, alert settings %v, %v, %s and %d",
cfg.AlertWebhookURL, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
}
}
func TestInvalidAlertSettingStopsTheStart(t *testing.T) {
t.Parallel()
wantStartStopped(t, []struct{ name, value string }{
{alertWebhookURL, "alerts.example/smallwebwaf"},
{alertWebhookURL, "ftp://alerts.example/"},
{alertWebhookURL, "https:///smallwebwaf"},
{alertWebhookURL, "https://user:password@alerts.example/"},
{alertWebhookURL, "https://alerts.example/#top"},
{alertWebhookURL, "https://alerts.example:0/"},
{alertWebhookURL, "https://alerts.example:65536/"},
{alertWebhookHeaders, "Authorization"},
{alertWebhookHeaders, "X Team:ops"},
{alertWebhookHeaders, ":ops"},
{alertWebhookHeaders, "X-Team:ops,"},
{alertWebhookHeaders, "X-Team:o\r\nps"},
{alertEvents, "bans"},
{alertEvents, "summary"},
{alertEvents, "ban,,file_error"},
{alertCooldown, "0"},
{alertCooldown, "soon"},
{alertMaxPerHour, "0"},
{alertMaxPerHour, "-1"},
{alertMaxPerHour, "1.5"},
})
}
func TestWebhookHeadersAreLoggedMaskedAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "Bearer 0123456789abcdef"
cfg := fromEnvironment(t, environment{
alertWebhookHeaders: "Authorization:" + secret + ",X-Team:ops",
})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) || strings.Contains(logged, "ops") ||
!strings.Contains(logged,
`"`+alertWebhookHeaders+`":"Authorization:********,X-Team:********"`) {
t.Errorf("the headers are not logged masked: %s", logged)
}
// An item that is not a header is named by its place, not shown.
_, err := config.FromEnvironment(environment{
alertWebhookHeaders: "X-Team:ops," + secret,
}.lookupEnv)
want := alertWebhookHeaders + ": item 2 is not a header name followed by : " +
"and the header's value, such as Authorization:Bearer <token>"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestWebhookURLIsLoggedWithoutItsPathOrQueryAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "T0123/B4567/abcdef"
for value, want := range map[string]string{
"https://hooks.example/services/" + secret: "https://hooks.example/********",
"https://hooks.example:8443?token=" + secret: "https://hooks.example:8443/********",
"http://[2001:db8::1]:8080": "http://[2001:db8::1]:8080",
} {
cfg := fromEnvironment(t, environment{alertWebhookURL: value})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) ||
!strings.Contains(logged, `"`+alertWebhookURL+`":"`+want+`"`) {
t.Errorf("%s is not logged as %s: %s", value, want, logged)
}
}
// A value that is not such a URL is not shown either.
for _, value := range []string{
"ftp://hooks.example/services/" + secret,
"https://hooks.example/services/%zz" + secret,
} {
_, err := config.FromEnvironment(environment{alertWebhookURL: value}.lookupEnv)
want := alertWebhookURL + ": is not an http or https URL without a user or " +
"a fragment, such as https://alerts.example/smallwebwaf"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
}
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
t.Parallel()
@@ -298,10 +736,8 @@ func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
func TestInvalidValueStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct{ name, value string }{
{listenAddr, "8080"},
{listenAddr, ":http"},
{listenAddr, ":65536"},
wantStartStopped(t, []struct{ name, value string }{
{listenAddr, "8080"}, {listenAddr, ":http"}, {listenAddr, ":65536"},
{upstreamURL, "127.0.0.1:8081"},
{upstreamURL, "ftp://127.0.0.1:8081"},
{upstreamURL, "http://"},
@@ -319,8 +755,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowNets, "192.0.2.0/24,monitoring"},
{rateLimitExemptNets, "2001:db8::/129"},
{denyNets, "198.51.100.0/24,"},
{clientRequestTimeout, "60"},
{clientRequestTimeout, ""},
{clientRequestTimeout, "60"}, {clientRequestTimeout, ""},
{clientIdleTimeout, "0s"},
{clientIdleTimeout, "2 minutes"},
{clientResponseTimeout, "1y"},
@@ -337,8 +772,8 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{rateLimitPerMinute, "1K"},
{rateLimitPerHour, "0"},
{rateLimitPerHour, "1.5"},
{rateLimitPerDay, "-1"},
{rateLimitPerDay, "lots"},
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
{rateLimitExemptPaths, "/assets/,,/static/"},
{deniedCountries, "nk"},
{deniedCountries, "kp,,ir"},
{deniedCountries, "prk"},
@@ -351,17 +786,37 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowedCountries, "uk"},
{allowedCountries, "zz"},
{allowedCountries, "de,germany"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
{logRequestHeaders, "host"}, {logRequestHeaders, "accept,Host"},
{logRequestHeaders, "transfer-encoding"}, {logRequestHeaders, "TRANSFER-ENCODING"},
{rulesEnabled, "yes"}, {rulesEnabled, "True"},
})
}
func TestInvalidBanOrStateValueStopsTheStart(t *testing.T) {
t.Parallel()
wantStartStopped(t, []struct{ name, value string }{
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
{maxBanDuration, off}, {maxBanDuration, "1w"},
{maxBanDuration, off}, {maxBanDuration, "1w"}, {attackBanDuration, off},
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
} {
})
}
// wantStartStopped checks that each setting, set to its value, stops the
// start with an error that names the setting.
func wantStartStopped(t *testing.T, invalid []struct{ name, value string }) {
t.Helper()
for _, tc := range invalid {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
@@ -377,39 +832,195 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
}
}
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
func TestHostOrTransferEncodingStopsTheStart(t *testing.T) {
t.Parallel()
// Characters are counted, not bytes: each é takes two.
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
// Only Host's message points to the field host.
for value, want := range map[string]string{
"Host": `"Host" is taken out of every request by Go's HTTP server, ` +
"so it can never be logged; the request's host is the field host",
"transfer-encoding": `"transfer-encoding" is taken out of every ` +
"request by Go's HTTP server, so it can never be logged",
} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv)
want := metricsToken + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
_, err := config.FromEnvironment(environment{logRequestHeaders: value}.lookupEnv)
if err == nil || err.Error() != logRequestHeaders+": "+want {
t.Errorf("error %v, want %s: %s", err, logRequestHeaders, want)
}
})
}
}
func TestTokenIsLoggedMasked(t *testing.T) {
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{metricsToken: token})
// Characters are counted, not bytes: each é takes two.
for _, name := range []string{adminToken, metricsToken} {
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
t.Run(name+"="+value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{name: value}.lookupEnv)
want := name + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
}
func TestTokensAreReadAndLoggedMasked(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{adminToken: otherToken, metricsToken: token})
if cfg.AdminToken != otherToken || cfg.MetricsToken != token {
t.Errorf("admin token %q and metrics token %q, want %q and %q",
cfg.AdminToken, cfg.MetricsToken, otherToken, token)
}
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
if strings.Contains(out.String(), token) ||
!strings.Contains(out.String(), `"`+metricsToken+`":"********"`) {
t.Errorf("the token is not logged masked: %s", out.String())
logged := out.String()
if strings.Contains(logged, token) || strings.Contains(logged, otherToken) ||
!strings.Contains(logged, `"`+adminToken+`":"********"`) ||
!strings.Contains(logged, `"`+metricsToken+`":"********"`) {
t.Errorf("the tokens are not logged masked: %s", logged)
}
}
func TestSettingFromFileLosesOneNewlineAndNoMore(t *testing.T) {
t.Parallel()
for contents, want := range map[string]string{
token: token,
token + "\n": token,
token + "\n\n": token + "\n",
token + " \n": token + " ",
} {
cfg := fromEnvironment(t, environment{
metricsToken + "_FILE": writeFile(t, contents),
})
if cfg.MetricsToken != want {
t.Errorf("file holding %q gave %s %q, want %q", contents, metricsToken,
cfg.MetricsToken, want)
}
}
}
func TestSettingFromFileIsCheckedAsTheSettingItself(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
requestMaxBytes + "_FILE": writeFile(t, "lots\n"),
}.lookupEnv)
want := requestMaxBytes + `: "lots" is not a size such as 512K, 100M or 5G, or off`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
_, err = config.FromEnvironment(environment{
logRemoteURL: remoteURL,
instanceName: instance,
logRemoteAppName + "_FILE": writeFile(t, "my app\n"),
}.lookupEnv)
want = logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
`characters without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestSettingAndItsFileBothSetStopsTheStart(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
metricsToken: token,
metricsToken + "_FILE": writeFile(t, token),
}.lookupEnv)
want := metricsToken + ": is set, and so is " + metricsToken +
"_FILE; set only one of them"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestUnreadableSettingFileStopsTheStart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
for _, path := range []string{filepath.Join(dir, "missing"), dir} {
_, err := config.FromEnvironment(environment{metricsToken + "_FILE": path}.lookupEnv)
want := metricsToken + "_FILE: cannot be read: "
if err == nil || !strings.HasPrefix(err.Error(), want) {
t.Errorf("%s: error %v, want one starting %s", path, err, want)
}
}
}
func TestTokenFromFileIsLoggedMaskedWithTheFile(t *testing.T) {
t.Parallel()
path := writeFile(t, token+"\n")
cfg := fromEnvironment(t, environment{metricsToken + "_FILE": path})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
var line struct {
Settings map[string]string `json:"settings"`
}
err := json.Unmarshal(out.Bytes(), &line)
if err != nil {
t.Fatalf("decode %s: %v", out.Bytes(), err)
}
if strings.Contains(out.String(), token) ||
line.Settings[metricsToken] != "********" ||
line.Settings[metricsToken+"_FILE"] != path {
t.Errorf("the token is not logged masked, with its file %s: %s", path,
out.String())
}
}
func TestRemoteLogCAFileIsNotReadAsAFileInItsTurn(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
logRemoteTLSCAFile + "_FILE": writeFile(t, "/nonexistent/ca.pem\n"),
})
if cfg.LogRemoteTLSCAs != nil {
t.Errorf("%s_FILE gave certificates", logRemoteTLSCAFile)
}
}
// writeFile writes contents to a file in a directory of its own, removed
// when the test ends, and returns the file's path.
func writeFile(t *testing.T, contents string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "setting")
err := os.WriteFile(path, []byte(contents), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
return path
}
func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Parallel()
@@ -428,6 +1039,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",
@@ -447,19 +1060,36 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
rateLimitPerMinute: "1000",
rateLimitPerHour: "10000",
rateLimitPerDay: "50000",
rateLimitExemptPaths: "",
deniedCountries: "",
allowedCountries: "",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
maxBanDuration: "7d",
attackBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
adminToken: "",
metricsToken: "",
metricsTopN: "50",
instanceName: hostname,
logRequestHeaders: defaultLogRequestHeaders,
rulesDir: "/etc/smallwebwaf/rules.d",
rulesEnabled: "true",
logRemoteURL: "",
logRemoteTLSCAFile: "",
logRemoteBuffer: "10000",
logRemoteFacility: "local0",
logRemoteAppName: hostname,
alertWebhookURL: "",
alertWebhookHeaders: "",
alertEvents: defaultAlertEvents,
alertCooldown: defaultAlertCooldown,
alertMaxPerHour: "60",
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -489,8 +1119,8 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans, the state files and the
// metrics.
// wantBanSettings checks the settings for bans, the state files, the
// metrics and the rule files.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
@@ -498,11 +1128,17 @@ func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
got.LimitBanDuration != want.LimitBanDuration ||
got.LimitBanRepeatWindow != want.LimitBanRepeatWindow ||
got.MaxBanDuration != want.MaxBanDuration ||
got.AttackBanDuration != want.AttackBanDuration ||
got.MaxBans != want.MaxBans ||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
}
if got.RulesDir != want.RulesDir || got.RulesEnabled != want.RulesEnabled {
t.Errorf("rule files in %q, on: %t, want %q, %t",
got.RulesDir, got.RulesEnabled, want.RulesDir, want.RulesEnabled)
}
if got.StateDir != want.StateDir ||
got.StateWriteDelay != want.StateWriteDelay ||
got.StateCounterInterval != want.StateCounterInterval {
+13
View File
@@ -19,6 +19,7 @@ import (
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics"
)
@@ -68,6 +69,8 @@ type Params struct {
// Metrics count the requests to GeoJS, those that failed, and the
// clients that go without an answer.
Metrics *metrics.Metrics
// Alerts receive a source_failure alert each time GeoJS fails.
Alerts *alerts.Queue
}
// GeoJS looks up clients' countries through GeoJS. At most one request
@@ -78,6 +81,7 @@ type GeoJS struct {
now func() time.Time
processLog *slog.Logger
metrics *metrics.Metrics
alerts *alerts.Queue
// httpClient follows no redirect, so that visitors' addresses go to
// GeoJS alone: a redirect is a failure.
httpClient *http.Client
@@ -127,6 +131,7 @@ func New(params Params) *GeoJS {
now: params.Now,
processLog: params.ProcessLog,
metrics: params.Metrics,
alerts: params.Alerts,
httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
@@ -386,6 +391,14 @@ func (g *GeoJS) keep(
g.processLog.Warn("asking GeoJS failed",
"error", err.Error(), "asking_again_in", g.retryDelay.String())
g.alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs", "error": err.Error(),
"asking_again_in": g.retryDelay.String(),
},
})
return false
}
+61 -1
View File
@@ -6,6 +6,8 @@ import (
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"reflect"
"slices"
"strings"
"sync"
@@ -14,6 +16,7 @@ import (
"time"
"github.com/prometheus/client_golang/prometheus/testutil"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
)
@@ -197,6 +200,7 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1),
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(geojs)
@@ -211,6 +215,45 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
})
}
func TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g, queue := startWithAlerts()
clients := newClients()
geojs.set(failing)
wantCountry(t, g, clients(), "")
want := alerts.Alert{
Time: clock.Now(),
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs",
"error": "GeoJS answered 503 Service Unavailable",
"asking_again_in": "1s",
},
}
// The next failure, a second later, is a repeat within the
// cooldown.
clock.advance(time.Second)
wantCountry(t, g, clients(), "")
wantRequests(t, geojs, 2)
waiting := queue.Snapshot().Waiting
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
}
if queue.Suppressed() != 1 {
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
}
})
}
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
t.Parallel()
@@ -361,6 +404,7 @@ func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(&standIn{answers: failing})
@@ -524,17 +568,33 @@ func (c *testClock) advance(d time.Duration) {
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
// asking the stand-in by that clock.
func start() (*standIn, *testClock, *lookup.GeoJS) {
geojs, clock, g, _ := startWithAlerts()
return geojs, clock, g
}
// startWithAlerts is start, and returns the queue of the alerts GeoJS
// raises as well, for a webhook that is never sent them, with the default
// cooldown, by the same clock.
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
geojs := &standIn{}
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: clock.Now,
})
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1),
Alerts: queue,
})
g.SetTransport(geojs)
return geojs, clock, g
return geojs, clock, g, queue
}
// newClients returns what returns a new IPv4 client each time it is
+109 -9
View File
@@ -11,9 +11,12 @@ import (
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/collectors"
"github.com/prometheus/client_golang/prometheus/promhttp"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
@@ -30,7 +33,9 @@ type Metrics struct {
rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec
countries *countries
// ruleMatches are made by AddRules.
ruleMatches *prometheus.CounterVec
countries *countries
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
// that failed. GeoJSUnanswered are the requests whose client counted
@@ -135,19 +140,22 @@ func New(topN int) *Metrics {
// AddBansAndClients adds the metrics read from the ledger and the table
// of clients as the metrics are asked for: the bans made since the start,
// the bans active and permanent at now, and the clients in the table.
// by cause, the bans active and permanent at now, and the clients in the
// table.
func (m *Metrics) AddBansAndClients(
ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time,
) {
m.registry.MustRegister(
// Every ban smallwebwaf makes so far is for a broken limit.
prometheus.NewCounterFunc(prometheus.CounterOpts{
for _, cause := range []string{bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin} {
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_bans_made_total",
Help: "Bans made, by cause.",
ConstLabels: prometheus.Labels{"cause": "limit"},
ConstLabels: prometheus.Labels{"cause": cause},
}, func() float64 {
return float64(ledger.Made())
}),
return float64(ledger.Made(cause))
}))
}
m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_active_bans",
Help: "Bans active now, the permanent ones included.",
@@ -158,7 +166,7 @@ func (m *Metrics) AddBansAndClients(
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_permanent_bans",
Help: "Permanent bans.",
Help: "Permanent bans not lifted.",
}, func() float64 {
_, permanent := ledger.Count(now())
@@ -173,6 +181,92 @@ func (m *Metrics) AddBansAndClients(
)
}
// AddRules adds the metrics of the rule files: the requests that matched
// each rule, which RuleMatched counts, and the rules loaded from
// ruleFiles, read as the metrics are asked for. It is called once, before
// RuleMatched.
func (m *Metrics) AddRules(ruleFiles *rules.Files) {
m.ruleMatches = counterVec("smallwebwaf_rule_matches_total",
"Requests that matched a rule of the rule files, by its id and action.",
[]string{"rule_id", "action"})
m.registry.MustRegister(m.ruleMatches,
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_rules_loaded",
Help: "Rules loaded from the rule files.",
}, func() float64 {
return float64(ruleFiles.Len())
}))
}
// AddRemoteLog adds the metrics of sending the log lines to
// SWWAF_LOG_REMOTE_URL, read from remote as the metrics are asked for: the
// lines sent, those dropped, and those waiting in the buffer.
func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_remote_log_lines_sent_total",
Help: "Log lines sent to SWWAF_LOG_REMOTE_URL.",
}, func() float64 {
return float64(remote.Sent())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_remote_log_lines_dropped_total",
Help: "Log lines dropped: the oldest in a full buffer, and those " +
"whose sending failed.",
}, func() float64 {
return float64(remote.Dropped())
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_remote_log_buffer_depth",
Help: "Log lines in the buffer, waiting to be sent.",
}, func() float64 {
return float64(remote.Depth())
}),
)
}
// AddAlerts adds the metrics of the alerts sent to
// SWWAF_ALERT_WEBHOOK_URL, read from queue as the metrics are asked for,
// with the destination webhook: the alerts sent, the requests to the
// webhook that failed, and the alerts held back and dropped.
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
webhook := prometheus.Labels{"destination": "webhook"}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_sent_total",
Help: "Alerts the destination took.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Sent())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_failed_total",
Help: "Requests to the destination that failed.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Failed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_suppressed_total",
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Suppressed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_dropped_total",
Help: "Alerts dropped, the oldest first, from a full queue, and alerts " +
"given up as the destination refused them.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Dropped())
}),
)
}
// 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)
@@ -219,6 +313,12 @@ func (m *Metrics) RequestEnded(
}
}
// RuleMatched counts a request that matched the rule id, whose action is
// action.
func (m *Metrics) RuleMatched(id, action string) {
m.ruleMatches.WithLabelValues(id, action).Inc()
}
// StateFileWritten counts a write of the state file name, of size bytes,
// that ended with err.
func (m *Metrics) StateFileWritten(name string, size int, err error) {
+289 -7
View File
@@ -1,26 +1,60 @@
package proxy
import (
"bytes"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/netip"
"os"
"strings"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
// banBodyMaxBytes is the most of the body of a request to add a ban that
// is read; its three fields need far less.
const banBodyMaxBytes = 4 << 10
// permanent is how the log line and the ban endpoint name a ban that
// never ends.
const permanent = "permanent"
var (
errNotBanToAdd = errors.New(
"the body is not a JSON object of netblock, duration and reason")
errNotNetblock = errors.New(
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
errMappedNetblock = errors.New(
"is IPv4-mapped: give the IPv4 netblock, such as 203.0.113.0/24")
errZone = errors.New("has a zone, which a netblock cannot have")
errNotDuration = errors.New(
"is not a duration above zero, such as 1h or 7d, or permanent")
errNotAddress = errors.New("is not an address, such as 203.0.113.9")
)
// answerAdmin answers a request for smallwebwaf itself, under
// /_smallwebwaf/, once it has passed the checks: GET MetricsPath with
// SWWAF_METRICS_TOKEN gets the metrics, and without it is refused with
// 401. Any other request gets 404, as the metrics do while
// SWWAF_METRICS_TOKEN is unset.
// /_smallwebwaf/, once it has passed the checks. Each endpoint needs a
// token, sent as Authorization: Bearer <token>: the metrics
// SWWAF_METRICS_TOKEN, the others SWWAF_ADMIN_TOKEN. A request without
// it is refused with 401. An endpoint whose token is unset answers 404,
// as any other request under /_smallwebwaf/ does.
func (rq *request) answerAdmin() {
rq.line.Action = requestlog.ActionAdmin
rq.startClientResponseTimeout()
token := rq.h.config.MetricsToken
token, answer := rq.endpoint()
switch {
case token == "" || rq.in.Method != http.MethodGet || rq.in.URL.Path != MetricsPath:
case token == "":
http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound)
case !hasToken(rq.in, token):
rq.out.Header().Set("WWW-Authenticate", "Bearer")
@@ -29,7 +63,29 @@ func (rq *request) answerAdmin() {
action: requestlog.ActionAdmin,
})
default:
rq.h.metrics.ServeHTTP(rq.out, rq.in)
answer()
}
}
// endpoint returns the token the request's endpoint needs, and what
// answers the request there; "" when there is no such endpoint.
func (rq *request) endpoint() (string, func()) {
cfg := rq.h.config
method, path := rq.in.Method, rq.in.URL.Path
switch {
case method == http.MethodGet && path == MetricsPath:
return cfg.MetricsToken, func() { rq.h.metrics.ServeHTTP(rq.out, rq.in) }
case method == http.MethodGet && path == BansPath:
return cfg.AdminToken, rq.listBans
case method == http.MethodPost && path == BansPath:
return cfg.AdminToken, rq.addBan
case method == http.MethodDelete && strings.HasPrefix(path, BansPath+"/"):
return cfg.AdminToken, rq.liftBans
case method == http.MethodGet && strings.HasPrefix(path, ClientsPath):
return cfg.AdminToken, rq.showClient
default:
return "", nil
}
}
@@ -41,3 +97,229 @@ func hasToken(r *http.Request, token string) bool {
return strings.EqualFold(scheme, "Bearer") &&
subtle.ConstantTimeCompare([]byte(sent), []byte(token)) == 1
}
// listBans answers GET BansPath with every ban held.
func (rq *request) listBans() {
rq.answerBans(rq.h.ledger.Snapshot())
}
// banToAdd is the body of POST BansPath.
type banToAdd struct {
// Netblock is a netblock, or a client's address, which stands for the
// netblock a ban on that client covers.
Netblock string `json:"netblock"`
// Duration is how long the ban lasts, as a setting gives a duration,
// or permanent.
Duration string `json:"duration"`
Reason string `json:"reason"`
}
// addBan answers POST BansPath: it bans the netblock the body names, as
// an admin, from now for the duration the body gives, with its reason,
// and answers with that ban.
func (rq *request) addBan() {
// The body must arrive within SWWAF_CLIENT_REQUEST_TIMEOUT, as any
// other request's must.
rq.stopReadingBody(rq.clientRequestDeadline())
toAdd, err := rq.readBanToAdd()
if refused := rq.refused.Load(); refused != nil {
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
return
}
if errors.Is(err, os.ErrDeadlineExceeded) {
rq.answer(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
})
return
}
var (
netblock netip.Prefix
expires time.Time
now = rq.h.now()
)
if err == nil {
netblock, err = rq.h.banNetblock(toAdd.Netblock)
}
if err == nil {
expires, err = expiry(toAdd.Duration, now)
}
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
ban := rq.h.ledger.BanForAdmin(netblock, now, expires, toAdd.Reason)
rq.answerBans([]bans.Ban{ban})
}
// readBanToAdd reads the body of POST BansPath: a JSON object with
// nothing but whitespace after it, in at most banBodyMaxBytes.
func (rq *request) readBanToAdd() (banToAdd, error) {
var body io.ReadCloser = http.NoBody
if rq.body != nil {
body = rq.body
}
data, err := io.ReadAll(http.MaxBytesReader(nil, body, banBodyMaxBytes))
if err != nil {
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
}
var toAdd banToAdd
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
err = decoder.Decode(&toAdd)
if err != nil {
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
}
// Token returns io.EOF only when nothing but whitespace is left.
_, err = decoder.Token()
if !errors.Is(err, io.EOF) {
return banToAdd{}, fmt.Errorf("%w: more follows the object", errNotBanToAdd)
}
return toAdd, nil
}
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a
// client's address, which stands for the netblock a ban on that client
// covers. An IPv4-mapped netblock, such as ::ffff:203.0.113.0/120, is
// refused, since a client's address is looked up as IPv4 and a ban on it
// would refuse nothing, and so is a value with a zone.
func (h *handler) banNetblock(value string) (netip.Prefix, error) {
netblock, err := netip.ParsePrefix(value)
if err == nil {
if netblock.Addr().Is4In6() {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errMappedNetblock)
}
return netblock, nil
}
// ParsePrefix refuses a zone, but ParseAddr reads the /48 of
// 2001:db8::1%x/48 as part of the zone.
addr, err := netip.ParseAddr(value)
if err != nil {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
}
if addr.Zone() != "" {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errZone)
}
return h.netblock(addr), nil
}
// expiry returns when a ban made at now for duration ends: duration
// later, for a duration as a setting gives one, or zero for permanent.
func expiry(duration string, now time.Time) (time.Time, error) {
if duration == permanent {
return time.Time{}, nil
}
length, err := config.ParseDurationNotOff(duration)
if err != nil {
return time.Time{}, fmt.Errorf("duration %q %w", duration, errNotDuration)
}
return now.Add(length), nil
}
// liftBans answers DELETE BansPath/<client>: it lifts every ban active on
// a netblock the client's address is in, and answers with those bans, or
// with 404 when none is active.
func (rq *request) liftBans() {
client, err := pathAddress(rq.in.URL.Path, BansPath+"/")
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
lifted := rq.h.ledger.Lift(client, rq.h.now())
if len(lifted) == 0 {
http.Error(rq.out, "no ban is active on "+client.String(), http.StatusNotFound)
return
}
rq.answerBans(lifted)
}
// clientAnswer is the answer to GET ClientsPath<ip>: the client the
// address is, as clients.json holds it, or null when the table of
// clients does not hold it, and the bans on each netblock the address is
// in, as bans.json lists them.
type clientAnswer struct {
Client *ratelimit.Client `json:"client"`
Bans []state.BanEntry `json:"bans"`
}
// showClient answers GET ClientsPath<ip> with what smallwebwaf knows of
// the client: its counters, its history, which holds its country as last
// looked up and its offences, and its bans with their notes.
func (rq *request) showClient() {
addr, err := pathAddress(rq.in.URL.Path, ClientsPath)
if err != nil {
http.Error(rq.out, err.Error(), http.StatusBadRequest)
return
}
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
client, seen := rq.h.limiter.Client(clientGroup(addr))
if seen {
answer.Client = &client
}
rq.answerJSON(answer)
}
// pathAddress reads the client's address that follows prefix in path.
func pathAddress(path, prefix string) (netip.Addr, error) {
value := strings.TrimPrefix(path, prefix)
addr, err := netip.ParseAddr(value)
if err != nil {
return netip.Addr{}, fmt.Errorf("%q %w", value, errNotAddress)
}
return addr.Unmap(), nil
}
// answerBans answers with held under bans, as bans.json lists them.
func (rq *request) answerBans(held []bans.Ban) {
rq.answerJSON(struct {
Bans []state.BanEntry `json:"bans"`
}{state.BanEntries(held)})
}
// answerJSON answers with value as indented JSON.
func (rq *request) answerJSON(value any) {
body, err := json.MarshalIndent(value, "", " ")
if err != nil {
rq.h.processLog.Error("encoding an answer failed", "error", err.Error())
http.Error(rq.out, http.StatusText(http.StatusInternalServerError),
http.StatusInternalServerError)
return
}
rq.out.Header().Set("Content-Type", "application/json")
_, _ = rq.out.Write(append(body, '\n'))
}
+535
View File
@@ -0,0 +1,535 @@
package proxy_test
import (
"encoding/json"
"net/http"
"net/netip"
"slices"
"strconv"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
const (
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set, and adminBearer
// how a request carries it.
adminSecret = "fedcba9876543210fedcba9876543210"
adminBearer = "Bearer " + adminSecret
// adminClient is the client the tests' admin sends its requests from.
adminClient = "192.0.2.10"
// banOtherClient is the body of a request to ban otherClient for an
// hour.
banOtherClient = `{"netblock": "` + otherClient + `", "duration": "1h", ` +
`"reason": "probes for logins"}`
)
func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
t.Parallel()
// The metrics token is set, and opens none of them.
s, clk, server := startWithClock(t, "", map[string]string{metricsToken: token})
server.Ledger.BanForLimit(netip.MustParsePrefix(otherClient+"/32"), clk.Now(),
bans.Notes{})
before := server.Ledger.Snapshot()
// An empty token does not match the unset one either.
for _, authorization := range []string{adminBearer, bearer, "Bearer ", ""} {
for _, e := range adminEndpoints() {
s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
http.StatusNotFound, requestlog.ActionAdmin)
}
}
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before)
}
}
func TestAdminEndpointsNeedTheAdminToken(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
metricsToken: token,
})
// Listing the bans, banning otherClient, lifting that ban, and asking
// about otherClient, in that order. Without the admin token, with the
// metrics token, or with one that differs, each is refused, and
// changes nothing; with the admin token, it is answered.
for _, e := range adminEndpoints() {
before := server.Ledger.Snapshot()
for _, authorization := range []string{
"", bearer, "Bearer " + strings.ToUpper(adminSecret), "Basic " + adminSecret,
} {
got := s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
http.StatusUnauthorized, requestlog.ActionAdmin)
if got.header.Get("WWW-Authenticate") != "Bearer" {
t.Errorf("%s %s with %q was answered without WWW-Authenticate: Bearer",
e.method, e.path, authorization)
}
}
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
e.method, e.path, after, before)
}
got := s.admin(e.method, e.path, e.body, http.StatusOK)
if got.header.Get("Content-Type") != "application/json" {
t.Errorf("%s %s answered %q", e.method, e.path, got.header.Get("Content-Type"))
}
}
// Any other request under /_smallwebwaf/ is not found.
for _, e := range []adminEndpoint{
{http.MethodPut, proxy.BansPath, banOtherClient},
{http.MethodDelete, proxy.BansPath, ""},
{http.MethodGet, proxy.BansPath + "/" + otherClient, ""},
{http.MethodPost, proxy.ClientsPath + otherClient, ""},
{http.MethodGet, strings.TrimSuffix(proxy.ClientsPath, "/"), ""},
} {
s.admin(e.method, e.path, e.body, http.StatusNotFound)
}
}
func TestBanAddedListedAndLiftedThroughTheEndpoints(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
banScopeV4Prefix: "24",
})
// A ban on otherClient bans the /24 a ban on that client covers, so it
// refuses client too, for an hour.
start := clk.Now()
expires := start.Add(time.Hour)
want := state.BanEntry{
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Start: start,
Expires: &expires,
Cause: bans.CauseAdmin,
Reason: "probes for logins",
}
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK),
want)
line := s.get(client, http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != requestlog.FormatTime(expires) {
t.Errorf("the ban ends at %s, want %s", line.BanExpires, expires)
}
// Its notes count the request it refused.
want.Notes.Requests, want.Notes.Refused = 1, 1
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
// Ten minutes on, lifting the bans on client lifts that one, which is
// kept, marked lifted.
clk.advance(10 * time.Minute)
lifted := clk.Now()
want.Lifted = &lifted
wantBans(t, s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "",
http.StatusOK), want)
s.get(client, http.StatusOK, requestlog.ActionForward)
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
// No ban on it is active any more.
s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "", http.StatusNotFound)
}
func TestBanToAddGivesItsNetblockAndDuration(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
banScopeV4Prefix: "24",
})
start := clk.Now()
for _, tc := range []struct {
netblock, duration string
want string
length time.Duration // 0 for a permanent ban
}{
// An address stands for the netblock a ban on that client covers.
{client, "7d", "203.0.113.0/24", 7 * 24 * time.Hour},
{"::ffff:198.51.100.7", "90m", "198.51.100.0/24", 90 * time.Minute},
{"2001:db8:5::1", "permanent", "2001:db8:5::/64", 0},
// A netblock stands for itself, its bits past its length cleared.
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
} {
// Whitespace may follow the object.
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` +
"\r\n"
want := state.BanEntry{
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
}
if tc.length != 0 {
expires := start.Add(tc.length)
want.Expires = &expires
}
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, body, http.StatusOK), want)
}
}
func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{adminToken: adminSecret})
for _, tc := range []struct{ body, want string }{
{"", "the body is not a JSON object of netblock, duration and reason: EOF"},
{"netblock=203.0.113.9", "the body is not a JSON object"},
{
`{"netblock": "203.0.113.9", "duration": "1h", "until": "2027"}`,
`unknown field "until"`,
},
{
`{"netblock": "203.0.113", "duration": "1h"}`,
`netblock "203.0.113" is not an address or a netblock`,
},
// A client's address is looked up as IPv4, so a ban on an
// IPv4-mapped netblock would refuse nothing.
{
`{"netblock": "::ffff:203.0.113.0/120", "duration": "1h"}`,
`netblock "::ffff:203.0.113.0/120" is IPv4-mapped`,
},
// Read as an address, its zone would be "x/48", and its ban on the
// /64 around it.
{
`{"netblock": "2001:db8::1%x/48", "duration": "1h"}`,
`netblock "2001:db8::1%x/48" has a zone`,
},
{
`{"netblock": "fe80::1%eth0", "duration": "1h"}`,
`netblock "fe80::1%eth0" has a zone`,
},
// Anything but whitespace after the object.
{
`{"netblock": "203.0.113.9", "duration": "1h"}` +
`{"netblock": "198.51.100.0/24", "duration": "1h"}`,
"more follows the object",
},
{`{"netblock": "203.0.113.9", "duration": "1h"} x`, "more follows the object"},
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
{
`{"netblock": "203.0.113.9", "duration": "off"}`,
`duration "off" is not a duration above zero`,
},
{
`{"netblock": "203.0.113.9", "duration": "0s"}`,
`duration "0s" is not a duration above zero`,
},
{
`{"netblock": "203.0.113.9", "duration": "forever"}`,
`duration "forever" is not a duration above zero, such as 1h or 7d, ` +
`or permanent`,
},
// Over the 4 KiB read of a body, even when the object comes first.
{
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
strings.Repeat("x", 4<<10) + `"}`,
"request body too large",
},
{
`{"netblock": "203.0.113.9", "duration": "1h"}` + strings.Repeat(" ", 4<<10),
"request body too large",
},
} {
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
if !strings.Contains(string(got.body), tc.want) {
t.Errorf("%.80s was answered %q, want it to say %q", tc.body, got.body, tc.want)
}
}
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestBanToAddOverTheRequestSizeLimitIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
requestMaxBytes: "16",
})
// Sent in a chunk, its length is not announced, so that it is found
// over SWWAF_REQUEST_MAX_BYTES only as it is read.
chunk := `{"netblock": "203.0.113.9", "duration": "1h"}`
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
http.MethodPost, proxy.BansPath,
strconv.FormatInt(int64(len(chunk)), 16)+"\r\n"+chunk+"\r\n0\r\n\r\n",
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestBanToAddSlowerThanTheClientRequestTimeoutIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
metricsToken: token,
clientRequestTimeout: shortTimeoutSetting,
})
// The chunk announces 256 bytes and the rest of it never comes, so only
// the timeout ends the wait. A hold-up of the test process can only
// make the answer later, so the time is checked only for not being
// shorter than the timeout.
start := time.Now()
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
http.MethodPost, proxy.BansPath, "100\r\n"+`{"netblock": "203.0.113.9", `,
http.StatusRequestTimeout, requestlog.ActionTimedOut)
if took := time.Since(start); took < shortTimeout {
t.Errorf("answered after %s, before the timeout of %s ran out", took, shortTimeout)
}
wantLimitHits(t, s.addr, clientRequestTimeout, 1)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
rateLimitPerMinute: "2",
rateLimitExemptNets: adminClient,
})
start := clk.Now()
// Two of otherClient's requests are let through; the third breaks the
// limit of two a minute, and bans it.
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited)
// Asked about by its address in IPv6 form too.
for _, addr := range []string{otherClient, "::ffff:" + otherClient} {
var got struct {
Client *ratelimit.Client `json:"client"`
Bans []state.BanEntry `json:"bans"`
}
decode(t, s.admin(http.MethodGet, proxy.ClientsPath+addr, "", http.StatusOK), &got)
if got.Client == nil {
t.Fatalf("%s: no client", addr)
}
history := got.Client.History
if got.Client.Client != netip.MustParsePrefix(otherClient+"/32") ||
history.Requests != 3 || history.Forwarded != 2 || history.Refused != 1 ||
history.Offences.Limit != 1 || !history.FirstSeen.Equal(start) {
t.Errorf("%s: client %+v", addr, got.Client)
}
if len(got.Bans) != 1 || got.Bans[0].Cause != bans.CauseLimit ||
got.Bans[0].Reason != "requests per minute over the limit of 2" ||
got.Bans[0].Notes.Count != 3 {
t.Errorf("%s: bans %+v, want the one for the broken limit", addr, got.Bans)
}
}
// Of an address no request came from and no ban covers, nothing is
// known.
got := s.admin(http.MethodGet, proxy.ClientsPath+"198.51.100.99", "", http.StatusOK)
if string(got.body) != "{\n \"client\": null,\n \"bans\": []\n}\n" {
t.Errorf("an unknown client is answered\n%s", got.body)
}
s.admin(http.MethodGet, proxy.ClientsPath+"203.0.113", "", http.StatusBadRequest)
s.admin(http.MethodDelete, proxy.BansPath+"/203.0.113.0/24", "",
http.StatusBadRequest)
}
func TestBannedClientIsRefusedAtTheEndpointsEvenWithTheToken(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{adminToken: adminSecret})
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
// otherClient cannot lift its own ban either.
for _, e := range adminEndpoints() {
s.adminRequest(otherClient, adminBearer, e.method, e.path, e.body,
http.StatusForbidden, requestlog.ActionBanned)
}
}
func TestAdminRequestsCountTowardTheLimits(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
rateLimitPerMinute: "2",
})
// A request refused for a missing token and one answered count toward
// the limit of two a minute, so the next breaks it.
s.adminRequest(client, "", http.MethodGet, proxy.BansPath, "",
http.StatusUnauthorized, requestlog.ActionAdmin)
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusOK, requestlog.ActionAdmin)
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusForbidden, requestlog.ActionRateLimited)
}
func TestClientInAllowNetsSkipsTheChecksButNeedsTheToken(t *testing.T) {
t.Parallel()
const allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
s, clk, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
allowNets: allowed,
rateLimitPerMinute: "1",
})
// A ban on it refuses nothing, and its requests are not counted.
server.Ledger.BanForAdmin(netip.MustParsePrefix(allowed+"/32"), clk.Now(),
time.Time{}, "")
for range 2 {
s.adminRequest(allowed, "", http.MethodGet, proxy.BansPath, "",
http.StatusUnauthorized, requestlog.ActionAdmin)
s.adminRequest(allowed, adminBearer, http.MethodGet, proxy.BansPath, "",
http.StatusOK, requestlog.ActionAdmin)
}
}
func TestAdminEndpointsNeedTheTokenInObserveMode(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
mode: observe,
})
for _, e := range adminEndpoints() {
s.adminRequest(adminClient, "", e.method, e.path, e.body,
http.StatusUnauthorized, requestlog.ActionAdmin)
}
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
// adminEndpoint is a request to an endpoint SWWAF_ADMIN_TOKEN opens.
type adminEndpoint struct {
method, path, body string
}
// adminEndpoints returns a request to each endpoint SWWAF_ADMIN_TOKEN
// opens: listing the bans, banning otherClient for an hour, lifting the
// bans on otherClient, and asking about otherClient.
func adminEndpoints() []adminEndpoint {
return []adminEndpoint{
{http.MethodGet, proxy.BansPath, ""},
{http.MethodPost, proxy.BansPath, banOtherClient},
{http.MethodDelete, proxy.BansPath + "/" + otherClient, ""},
{http.MethodGet, proxy.ClientsPath + otherClient, ""},
}
}
// admin sends a request with method for path, with body, from
// adminClient, with the admin token, and checks that it is answered with
// status, its log line's action admin. It returns the answer.
func (s *sender) admin(method, path, body string, status int) answer {
s.t.Helper()
return s.adminRequest(adminClient, adminBearer, method, path, body, status,
requestlog.ActionAdmin)
}
// adminRequest sends a request with method for path, with body, from the
// client at from, with authorization as its Authorization header unless
// it is "", and checks its answer's status and its log line's action, as
// request does. authorization may end in more header lines. A body that
// is not "" has its length announced, unless authorization names
// Transfer-Encoding. It returns the answer.
func (s *sender) adminRequest(
from, authorization, method, path, body string, status int, action string,
) answer {
s.t.Helper()
var header []string
if authorization != "" {
header = append(header, "Authorization: "+authorization)
}
if body != "" && !strings.Contains(authorization, "Transfer-Encoding") {
header = append(header, "Content-Length: "+strconv.Itoa(len(body)))
}
_, got := s.requestWithBody(method, from, path, strings.Join(header, "\r\n"),
body, status, action)
return got
}
// wantBans checks that a ban endpoint answered with want, and no other
// ban.
func wantBans(t *testing.T, got answer, want ...state.BanEntry) {
t.Helper()
var decoded struct {
Bans []state.BanEntry `json:"bans"`
}
decode(t, got, &decoded)
gotJSON, err := json.Marshal(decoded.Bans)
if err != nil {
t.Fatalf("encode %+v: %v", decoded.Bans, err)
}
wantJSON, err := json.Marshal(want)
if err != nil {
t.Fatalf("encode %+v: %v", want, err)
}
if string(gotJSON) != string(wantJSON) {
t.Errorf("bans\n%s\nwant\n%s", gotJSON, wantJSON)
}
}
// decode reads the JSON answer of an endpoint into value.
func decode(t *testing.T, got answer, value any) {
t.Helper()
err := json.Unmarshal(got.body, value)
if err != nil {
t.Fatalf("decode %s: %v", got.body, err)
}
}
+263
View File
@@ -0,0 +1,263 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"reflect"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
// alertInstance is the instance every alert of these tests gives.
alertInstance = "fsn1app1/gitea"
)
func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rateLimitPerMinute: "1",
banScopeV4Prefix: "24",
})
start := clk.Now()
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
netblock := netip.MustParsePrefix("203.0.113.0/24")
ban := server.Ledger.Bans(netblock)[0]
// A request refused under the ban raises no other alert.
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netblock, Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1", Notes: ban.Notes,
}, requestlog.FormatTime(start.Add(time.Hour))))
if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" {
t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes)
}
}
func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
netblock := netip.MustParsePrefix(client + "/32")
other := netip.MustParsePrefix(otherClient + "/32")
// The probe bans the client for seven days, and its next request makes
// the ban permanent. The request after that changes nothing.
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
attackBan := server.Ledger.Bans(netblock)[0]
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
permanentBan := server.Ledger.Bans(netblock)[0]
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
// Another client's probe after its first ban has run out without a
// request makes a permanent ban at once.
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
clk.advance(7 * 24 * time.Hour)
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
otherBans := server.Ledger.Bans(other)
wantAlerts(t, queue,
attackAlert(alerts.EventBan, start, client, attackBan,
requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute), client,
permanentBan, "permanent"),
attackAlert(alerts.EventBan, start.Add(time.Minute), otherClient, otherBans[0],
requestlog.FormatTime(start.Add(time.Minute+7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute+7*24*time.Hour),
otherClient, otherBans[1], "permanent"),
)
}
func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
// A ban for a clear sign of attack, which a request under it would make
// permanent.
group := netip.MustParsePrefix(ipv6Group)
attackBan, _ := server.Ledger.BanForAttack(group, start, bans.Notes{RuleID: "probe"})
// The third request breaks the limit, and so does the fourth, within the
// cooldown, which raises nothing. The probe is a clear sign of attack.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
// No ban is made, and none made permanent.
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan ||
line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+
"for the attack alone, as it was", held, line.BanExpires)
}
waiting := queue.Snapshot().Waiting
if len(waiting) != 3 || queue.Suppressed() != 0 {
t.Fatalf("%d alerts wait and %d are held back, want 3 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
limitNotes, _ := waiting[0].Detail["notes"].(bans.Notes)
attackNotes, _ := waiting[1].Detail["notes"].(bans.Notes)
if limitNotes.Limit != 2 || limitNotes.Request.Path != "/" ||
attackNotes.Request.Path != "/.env" {
t.Errorf("the notes are %+v and %+v, want those of the broken limit and "+
"of the probe", limitNotes, attackNotes)
}
// Each alert is the one enforce mode would have raised, with mode
// observe in its detail.
want := []alerts.Alert{
banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 2", Notes: limitNotes,
}, requestlog.FormatTime(start.Add(time.Hour))),
attackAlert(alerts.EventBan, start, otherClient, bans.Ban{
Netblock: netip.MustParsePrefix(otherClient + "/32"), Notes: attackNotes,
}, requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start, ipv6Client, attackBan, permanent),
}
for _, alert := range want {
alert.Detail["mode"] = observe
}
wantAlerts(t, queue, want...)
}
func TestObserveModeWorksOutABanOnlyWhenItsAlertWouldBeSent(t *testing.T) {
t.Parallel()
s, _, _, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
alertMaxPerHour: "2",
})
// The client's third request breaks the limit, and raises the first
// alert of the hour. Its fourth is within the cooldown.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
// The other client's first probe raises the second. Its second probe is
// within the cooldown.
for range 2 {
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
}
// The IPv6 client's third request breaks the limit past the two alerts
// an hour.
for range 3 {
s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
}
// Had the ban been worked out for any of the requests within the
// cooldown or past the two an hour, its alert would have been raised,
// held back and counted.
if waiting := queue.Snapshot().Waiting; len(waiting) != 2 || queue.Suppressed() != 0 {
t.Errorf("%d alerts wait and %d are held back, want 2 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
}
// startWithAlerts is startWithClock with alerts to a webhook, which is
// never sent them, and returns the queue they wait in as well.
func startWithAlerts(
t *testing.T, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{
trustedProxies: trustLocalhost,
alertWebhookURL: "https://alerts.example/smallwebwaf",
instanceName: alertInstance,
}
maps.Copy(settings, env)
addr, out, server, queue := startProxyWithAlerts(t, app.URL, "", clk.Now, settings)
return &sender{t: t, addr: addr, out: out}, clk, server, queue
}
// banAlert returns the alert for event, raised by a request from client at
// the time raised, for ban, with its netblock, cause, reason and notes,
// which ends at expires, as the log line gives it.
func banAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return alerts.Alert{
Instance: alertInstance,
Time: raised,
Event: event,
Client: netip.MustParseAddr(client),
Netblock: ban.Netblock,
Reason: ban.Reason,
Detail: map[string]any{
"cause": ban.Cause, "ban_expires": expires, "notes": ban.Notes,
},
}
}
// attackAlert is banAlert for a ban for the probe rule of testRules, with
// the netblock and the notes of ban.
func attackAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return banAlert(event, raised, client, bans.Ban{
Netblock: ban.Netblock, Cause: bans.CauseAttack, Reason: "matched the rule probe",
Notes: ban.Notes,
}, expires)
}
// wantAlerts checks the alerts waiting in queue, in order.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper()
got := queue.Snapshot().Waiting
if len(got) != len(want) {
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
}
for i := range want {
if !reflect.DeepEqual(got[i], want[i]) {
t.Errorf("alert %d is\n%+v\nwant\n%+v", i, got[i], want[i])
}
}
}
+151 -31
View File
@@ -4,8 +4,10 @@ import (
"net/netip"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged
@@ -15,29 +17,40 @@ func (rq *request) banResponse(action string) *refusal {
}
// banned reports whether a ban on a netblock the client is in covers the
// request at now, and notes for the log line when that ban ends.
// request at now, and notes for the log line when that ban ends. A
// request that makes the ban permanent, or in observe mode would have,
// raises the alert for it.
func (rq *request) banned(now time.Time) bool {
check := rq.h.ledger.Check
if rq.h.config.Observe {
check = rq.h.ledger.Find // in observe mode the ban refuses nothing
check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is
}
ban, banned := check(rq.client, now)
ban, banned, madePermanent := check(rq.client, now)
if banned {
rq.line.BanExpires = banExpires(ban)
}
if madePermanent {
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
rq.alertBan(ban)
}
return banned
}
// limitBroken counts the request for the rate limits at now, and reports
// whether it takes the client over one. In enforce mode such a request
// bans the client's netblock, and sets the client's counters back to
// zero; in observe mode it does neither.
// limitBroken counts the request for the rate limits at now, notes the
// client's counts for the log line, and reports whether the request takes
// the client over a limit. In enforce mode such a request bans the
// client's netblock, and sets the client's counters back to zero; in
// observe mode it does neither, and raises the alert for the ban it would
// have made, if that alert would be sent.
func (rq *request) limitBroken(now time.Time) bool {
group := clientGroup(rq.client)
hit, over := rq.h.limiter.Count(group, now)
counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts
if !over {
return false
}
@@ -45,40 +58,147 @@ func (rq *request) limitBroken(now time.Time) bool {
rq.line.LimitHit = hit.Window
rq.line.Offence = requestlog.OffenceLimit
if rq.h.config.Observe {
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
return true
}
netblock := rq.netblock()
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
},
// The histories count this request only once it has ended.
Requests: rq.h.limiter.Requests(netblock) + 1,
})
notes := bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return true
}
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
rq.h.limiter.Reset(group)
rq.line.BanExpires = banExpires(ban)
if made {
rq.alertBan(ban)
}
return true
}
// netblock is the netblock a ban on the client covers: its IPv4 address,
// banForAttack bans the client's netblock at now for a clear sign of
// attack, the match of rule, a ban rule. In observe mode it makes no ban,
// and raises the alert for the ban it would have made, if that alert
// would be sent.
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
return
}
notes := bans.Notes{
Country: rq.line.Country,
RuleID: rule.ID,
Target: rule.Target,
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return
}
ban, made := rq.h.ledger.BanForAttack(netblock, now, notes)
rq.line.BanExpires = banExpires(ban)
if made {
rq.alertBan(ban)
}
}
// wouldAlertBan reports whether the alert for a ban on netblock for cause
// made at now would be sent. In observe mode the ban the request would
// have made is worked out only then, at most once per
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
// netblock's requests, which can mean going through every client.
func (rq *request) wouldAlertBan(
netblock netip.Prefix, now time.Time, cause string,
) bool {
event := alerts.EventBan
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
event = alerts.EventPermanentBan
}
return rq.h.alerts.WouldSend(event, netblock)
}
// alertBan raises the alert for ban, which the request made, or made
// permanent: permanent_ban for a permanent ban, ban for another. Its
// detail gives the ban's cause, when it ends, and its notes, and in
// observe mode, where ban is the ban that would have been made, or made
// permanent, mode, observe.
func (rq *request) alertBan(ban bans.Ban) {
event := alerts.EventBan
if ban.Permanent() {
event = alerts.EventPermanentBan
}
detail := map[string]any{
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
}
if rq.h.config.Observe {
detail["mode"] = "observe"
}
rq.h.alerts.Raise(alerts.Alert{
Event: event,
Client: rq.client,
Netblock: ban.Netblock,
Country: ban.Notes.Country,
Reason: ban.Reason,
Detail: detail,
})
}
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, or in
// observe mode as it would have been, as the notes of the ban it makes
// keep it.
func (rq *request) noted(now time.Time) bans.Request {
return bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
}
}
// netblockRequests is how many requests netblock has sent since it was
// first seen, this one included: the histories count it only once it has
// ended.
func (rq *request) netblockRequests(netblock netip.Prefix) int64 {
return rq.h.limiter.Requests(netblock) + 1
}
// netblock is the netblock a ban on client covers: its IPv4 address,
// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup
// counts it in.
func (rq *request) netblock() netip.Prefix {
addr := rq.client.Unmap()
func (h *handler) netblock(client netip.Addr) netip.Prefix {
addr := client.Unmap()
if addr.Is4() {
return netip.PrefixFrom(addr, rq.h.config.BanScopeV4Prefix).Masked()
return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked()
}
return clientGroup(addr)
@@ -88,7 +208,7 @@ func (rq *request) netblock() netip.Prefix {
// permanent.
func banExpires(ban bans.Ban) string {
if ban.Permanent() {
return "permanent"
return permanent
}
return requestlog.FormatTime(ban.Expires)
+22 -6
View File
@@ -278,6 +278,8 @@ func TestBanNotes(t *testing.T) {
Netblock: netblock,
Start: start,
Expires: start.Add(time.Hour),
Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1",
Notes: bans.Notes{
Country: "DE",
Limit: 1,
@@ -295,7 +297,7 @@ func TestBanNotes(t *testing.T) {
// refused under the ban.
Requests: 4,
Refused: 2,
EarlierBans: 0,
EarlierBans: bans.EarlierBans{},
},
}
@@ -312,8 +314,8 @@ func TestBanNotes(t *testing.T) {
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
got = ledger.Bans(netblock)
if len(got) != 2 || got[1].Notes.EarlierBans != 1 {
t.Errorf("bans %+v, want two, the second with one earlier ban", got)
if len(got) != 2 || got[1].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("bans %+v, want two, the second with one earlier ban for a limit", got)
}
}
@@ -415,14 +417,28 @@ func (s *sender) requestWithHeader(
) (logLine, string) {
s.t.Helper()
line, got := s.requestWithBody(http.MethodGet, from, path, header, "", status, action)
return line, string(got.body)
}
// requestWithBody is requestWithHeader for a request with method, whose
// body is sent as it is after the headers, header holding its
// Content-Length or Transfer-Encoding. header may hold several lines,
// separated by "\r\n". It returns the whole answer.
func (s *sender) requestWithBody(
method, from, path, header, body string, status int, action string,
) (logLine, answer) {
s.t.Helper()
if header != "" {
header += "\r\n"
}
conn := dial(s.t, s.addr)
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+
send(s.t, conn, method+" "+path+" HTTP/1.1\r\nHost: "+appHost+
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+
header+"\r\n")
header+"\r\n"+body)
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
@@ -451,5 +467,5 @@ func (s *sender) requestWithHeader(
s.sent++
wantLine(s.t, line, status, action)
return line, string(got.body)
return line, got
}
+28
View File
@@ -1,6 +1,7 @@
package proxy
import (
"crypto/rand"
"net/http"
"net/netip"
"slices"
@@ -48,6 +49,33 @@ func clientAddress(
return client
}
// requestIDHeader carries the request's id, from traefik and to the app.
const requestIDHeader = "X-Request-ID"
// requestID is the request's id: the one a trusted proxy sent, or a new
// random one. A peer outside the trusted proxies did not come through
// traefik, so the id it sends is its own claim, and is replaced.
func requestID(r *http.Request, peerTrusted bool) string {
id := r.Header.Get(requestIDHeader)
if !peerTrusted || id == "" {
id = rand.Text()
}
return id
}
// scheme is how the client reached traefik, as a trusted proxy says in
// X-Forwarded-Proto, or otherwise http, the only scheme smallwebwaf
// serves.
func scheme(r *http.Request, peerTrusted bool) string {
proto := r.Header.Get("X-Forwarded-Proto")
if !peerTrusted || proto == "" {
return "http"
}
return proto
}
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
const ipv6GroupPrefix = 64
+17 -13
View File
@@ -14,10 +14,14 @@ const (
appHost = "app.example"
// client is the client's address, as a proxy names it.
client = "203.0.113.9"
// forwardedFor is the header that lists the client and its proxies.
forwardedFor = "X-Forwarded-For"
// secure is the scheme a client reached traefik with.
// forwardedFor is the header that lists the client and its proxies,
// and forwardedProto the one that gives the scheme the client used.
forwardedFor = "X-Forwarded-For"
forwardedProto = "X-Forwarded-Proto"
// secure is the scheme a client reached traefik with, and plain the
// one smallwebwaf serves.
secure = "https"
plain = "http"
)
// appHeaders is what the app tells about the headers it received.
@@ -65,13 +69,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
func clientAddressCases() []clientAddressCase {
trusted := map[string]string{trustedProxies: trustLocalhost}
forged := http.Header{
forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"},
forwardedProto: {secure},
"X-Real-Ip": {client},
}
replaced := appHeaders{
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http",
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
}
return []clientAddressCase{{
@@ -87,10 +91,10 @@ func clientAddressCases() []clientAddressCase {
"outside the trusted proxies from the right",
env: trusted,
header: http.Header{
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost},
forwardedProto: {secure},
"X-Real-Ip": {client},
},
wantClient: client,
wantApp: appHeaders{
@@ -138,7 +142,7 @@ func requestWithHeaders(
Host: r.Host,
ForwardedFor: r.Header.Get(forwardedFor),
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
ForwardedProto: r.Header.Get("X-Forwarded-Proto"),
ForwardedProto: r.Header.Get(forwardedProto),
RealIP: r.Header.Get("X-Real-IP"),
})
})
+12 -3
View File
@@ -21,14 +21,18 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
// the last one would have it refused.
addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"})
const healthChecks = 3
const (
healthChecks = 3
contentType = "text/plain; charset=utf-8"
)
for range healthChecks {
got := get(t, addr, proxy.HealthPath)
wantStatus(t, got, http.StatusOK)
if string(got.body) != "ok\n" {
t.Errorf("health endpoint answered %q, want ok", got.body)
if string(got.body) != "ok\n" || got.header.Get("Content-Type") != contentType {
t.Errorf("health endpoint answered %q with Content-Type %q, want ok "+
"with %q", got.body, got.header.Get("Content-Type"), contentType)
}
}
@@ -37,6 +41,11 @@ func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) {
lines := out.requestLines(t, healthChecks+1)
for _, line := range lines[:healthChecks] {
wantLine(t, line, http.StatusOK, requestlog.ActionAdmin)
if line.ResponseContentType != contentType {
t.Errorf("health check's log line has response_content_type %q, "+
"want %q", line.ResponseContentType, contentType)
}
}
wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward)
+29
View File
@@ -12,6 +12,7 @@ import (
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
@@ -243,6 +244,34 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
}
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
s, clk, server := startWithClock(t, "", map[string]string{
metricsToken: token,
adminToken: adminSecret,
rateLimitExemptNets: scraper,
})
const admins = `smallwebwaf_bans_made_total{cause="admin"}`
wantMetric(t, s.scrape(scraper), admins, 0)
// As an admin's edit of bans.json that adds a ban is taken in.
server.Ledger.LoadEdit([]bans.Ban{{
Netblock: netip.MustParsePrefix(client + "/32"),
Start: clk.Now(),
}})
wantMetric(t, s.scrape(scraper), admins, 1)
// And a ban made through the endpoint.
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
wantMetric(t, s.scrape(scraper), admins, 2)
}
func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
t.Parallel()
+1
View File
@@ -95,6 +95,7 @@ func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
Netblock: netip.MustParsePrefix(otherClient + "/32"),
Start: clk.Now(),
Expires: clk.Now().Add(time.Hour),
Cause: bans.CauseAdmin,
}
server.Ledger.Load([]bans.Ban{kept})
+30 -15
View File
@@ -6,6 +6,8 @@ import (
"errors"
"io"
"net/http"
"os"
"reflect"
"slices"
"strings"
"sync/atomic"
@@ -14,6 +16,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
@@ -115,27 +118,34 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
}
}
// wantRequestFields checks the log line's fields about the request.
// wantRequestFields checks the log line's fields about the request. Its
// time, its id and its timings are checked only for being there.
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper()
want := requestlog.Line{
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery,
Protocol: "HTTP/1.1", Status: http.StatusTeapot,
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent),
hostname, _ := os.Hostname()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: hostname,
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent),
ResponseBytes: int64(received), UserAgent: "test-agent",
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal,
DurationUpstreamTotal: line.DurationUpstreamTotal,
}
if line.Line != want {
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
_, err := time.Parse(time.RFC3339, line.Time)
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 {
t.Errorf("log line has time %q and durations %v and %v",
line.Time, line.DurationTotal, line.DurationUpstreamTotal)
if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
t.Errorf("log line has time %q, request_id %q and durations %v and %v",
line.Time, line.RequestID, line.DurationTotal,
line.fields["duration_upstream_total"])
}
}
@@ -371,8 +381,13 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
addr, out := startProxy(t, "http://"+localhost+":1", nil)
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
wantLine(t, out.requestLine(t), http.StatusBadGateway,
requestlog.ActionUpstreamError)
line := out.requestLine(t)
wantLine(t, line, http.StatusBadGateway, requestlog.ActionUpstreamError)
// There never was a connection to the app, nor an answer from it.
wantTimings(t, line, "duration_total", "duration_checks",
"duration_upstream_total")
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
return line["type"] == "process" && line["msg"] == "request to the app failed"
+28
View File
@@ -11,12 +11,14 @@ import (
"strings"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// How smallwebwaf keeps connections to the app open between requests.
@@ -37,6 +39,14 @@ const HealthPath = "/_smallwebwaf/healthz"
// SWWAF_METRICS_TOKEN.
const MetricsPath = "/_smallwebwaf/metrics"
// BansPath is where an admin lists and adds bans, and, followed by / and
// a client's address, lifts them, with SWWAF_ADMIN_TOKEN.
const BansPath = "/_smallwebwaf/bans"
// ClientsPath is where an admin asks what smallwebwaf knows of a client,
// by the client's address after it, with SWWAF_ADMIN_TOKEN.
const ClientsPath = "/_smallwebwaf/clients/"
// Params are what New needs.
type Params struct {
Config *config.Config
@@ -51,6 +61,12 @@ type Params struct {
// limits, bans are made and run out, and GeoJS's answers are kept,
// normally time.Now in UTC, the time the state files give.
Now func() time.Time
// Rules are the rule files' rules, which each request is checked
// against.
Rules *rules.Files
// Alerts receive the alert for each ban the proxy makes or makes
// permanent, and for GeoJS failing.
Alerts *alerts.Queue
}
// Server is the server smallwebwaf runs, with the parts of the proxy
@@ -90,6 +106,7 @@ func New(params Params) *Server {
LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
MaxBanDuration: params.Config.MaxBanDuration,
AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans,
}),
geojs: lookup.New(lookup.Params{
@@ -97,9 +114,13 @@ func New(params Params) *Server {
Now: params.Now,
ProcessLog: params.ProcessLog,
Metrics: m,
Alerts: params.Alerts,
}),
rules: params.Rules,
alerts: params.Alerts,
}
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
m.AddRules(params.Rules)
return &Server{
Server: &http.Server{
@@ -134,6 +155,8 @@ type handler struct {
limiter *ratelimit.Limiter
ledger *bans.Ledger
geojs *lookup.GeoJS
rules *rules.Files
alerts *alerts.Queue
}
// newTransport returns what carries requests to the app. It never goes
@@ -160,6 +183,9 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// a health checker is never refused. It does not ask the app.
if r.Method == http.MethodGet && r.URL.Path == HealthPath {
rq.line.Action = requestlog.ActionAdmin
// Set here rather than left to Go's server, which would set it only
// after the log line has taken the response's headers.
rq.out.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = io.WriteString(rq.out, "ok\n")
return
@@ -169,6 +195,8 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
defer rq.addToHistory()
refused := rq.check(r.Context())
rq.checked = time.Now()
if refused != nil {
rq.answer(*refused)
+60 -5
View File
@@ -14,9 +14,11 @@ import (
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
const (
@@ -35,6 +37,10 @@ const (
// localhost is where every test server listens, and so the address
// smallwebwaf sees each test's requests come from.
localhost = "127.0.0.1"
// requestType is the type that marks a request log line.
requestType = "request"
// protocol is the protocol of every test's requests.
protocol = "HTTP/1.1"
)
// shortTimeoutSetting is shortTimeout as a setting's value.
@@ -59,6 +65,7 @@ const (
denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
@@ -67,6 +74,10 @@ const (
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
rulesDir = "SWWAF_RULES_DIR"
)
// output collects what smallwebwaf writes on stdout.
@@ -83,6 +94,14 @@ func (o *output) Write(p []byte) (int, error) {
return o.buf.Write(p)
}
// text returns everything written so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// lines returns every line written so far, decoded.
func (o *output) lines(t *testing.T) []map[string]any {
t.Helper()
@@ -122,7 +141,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
var found []logLine
for _, fields := range o.lines(t) {
if fields["type"] == "request" {
if fields["type"] == requestType {
found = append(found, decodeLine(t, fields))
}
}
@@ -197,14 +216,29 @@ func startProxyWithGeoJS(
}
// startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well.
// bans made by the time now tells, and returns the server as well. Unless
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *proxy.Server) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env)
return addr, out, server
}
// startProxyWithAlerts is startProxyWithClock, and returns the queue of
// the alerts the proxy raises as well, as the settings in env make it. No
// alert is sent from it: they wait in it, for the test to look at.
func startProxyWithAlerts(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *proxy.Server, *alerts.Queue) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()}
maps.Copy(settings, env)
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
@@ -217,12 +251,33 @@ func startProxyWithClock(
}
out := &output{}
processLog := requestlog.NewProcessLogger(out)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
})
if err != nil {
t.Fatalf("rule files: %v", err)
}
alertQueue := alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: out,
ProcessLog: requestlog.NewProcessLogger(out),
ProcessLog: processLog,
GeoJSURL: geojsURL,
Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
})
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
@@ -238,7 +293,7 @@ func startProxyWithClock(
_ = server.Close()
})
return listener.Addr().String(), out, server
return listener.Addr().String(), out, server, alertQueue
}
// newClient returns an HTTP client that sends requests as they are made,
+88
View File
@@ -5,6 +5,8 @@ import (
"sync/atomic"
"testing"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
@@ -69,3 +71,89 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
t.Errorf("the app was called %d times, want 4", calls.Load())
}
}
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
s, _, server := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/,/favicon.ico",
denyNets: denied,
deniedCountries: "kp",
})
// The answers are kept before the requests, so that none waits for
// GeoJS.
server.GeoJS.Load([]lookup.Answer{
keptAnswer(client, "DE"), keptAnswer(fromKP, "KP"),
})
// With a limit of one request a minute, the requests for paths under a
// prefix are not counted, so client's first request for / is within
// the limit; and once client has reached it, they are not refused.
s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
if line.LimitHit != "" || line.Counts != (ratelimit.Counts{}) {
t.Errorf("log line has limit_hit %q and counts %+v, want neither",
line.LimitHit, line.Counts)
}
// A path outside every prefix is counted: /assets is not under
// /assets/, and breaks the limit.
s.request(client, "/assets", http.StatusForbidden, requestlog.ActionRateLimited)
// A ban, SWWAF_DENY_NETS and the country lists still refuse a path
// under a prefix.
s.request(client, "/assets/app.js", http.StatusForbidden, requestlog.ActionBanned)
s.request(denied, "/assets/app.js", http.StatusForbidden, requestlog.ActionDenied)
s.request(fromKP, "/assets/app.js",
http.StatusForbidden, requestlog.ActionCountryDenied)
}
func TestRateLimitCountsPathsThatAreNotExempt(t *testing.T) {
t.Parallel()
for _, sent := range []string{
// A prefix matches only at the start of the path.
"/static/assets/app.js",
// A prefix matches the path as sent: a router that matches the
// path as received does not take /%61ssets/x for a path under
// /assets/.
"/%61ssets/x",
// .. once percent-decoded: an app may act on these as /login, the
// last as a path under /sneak/app/ or as /assets/x.
"/assets/../login",
"/assets/%2e%2e/login",
"/assets/..%2Flogin",
"/assets/..;/login",
"/sneak/app/src/branch/main/..%2F..%2F..%2F..%2F..%2F..%2Fassets/x",
// Not under /assets/ as sent: Go's router takes /assets%2Fx for one
// path segment, not a path under /assets/.
"/assets%2Fx",
"/assets%2fx",
// Under /assets/ as sent, but holding an encoded slash, in either
// case, or a backslash: never exempt, whatever the prefix.
"/assets/x%2Fy",
"/assets/x%2fy",
`/assets/x\y`,
} {
t.Run(sent, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/",
})
// Counted, the second request breaks the limit of one request
// a minute.
s.request(client, sent, http.StatusOK, requestlog.ActionForward)
s.request(client, sent, http.StatusForbidden, requestlog.ActionRateLimited)
})
}
}
+164 -37
View File
@@ -7,7 +7,10 @@ import (
"net/http/httptrace"
"net/http/httputil"
"net/netip"
"net/url"
"os"
"slices"
"strings"
"sync"
"sync/atomic"
"time"
@@ -46,7 +49,9 @@ type request struct {
peer netip.Addr
peerTrusted bool
start time.Time
// upstreamStart is when the request was handed to the app.
// checked is when the checks were done, and upstreamStart when the
// request was handed to the app.
checked time.Time
upstreamStart time.Time
// cancel ends the request to the app.
cancel context.CancelFunc
@@ -56,26 +61,34 @@ type request struct {
complete bool
// mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them from its own; once
// timersStopped is set, none of them acts any more.
// own, and the transport starts and stops them, and notes the times
// below, from its own; once timersStopped is set, none of the timeouts
// acts any more.
mu sync.Mutex
timersStopped bool
clientRequestTimer *time.Timer
upstreamRequestTimer *time.Timer
upstreamResponseTimer *time.Timer
// requestSent is when the app had been sent the whole request.
requestSent time.Time
// connected is when there was a connection to the app, requestSent
// when the app had been sent the whole request, and answerStarted
// when the first byte of its answer arrived.
connected time.Time
requestSent time.Time
answerStarted time.Time
}
// newRequest starts handling r: it notes the time, counts the request as
// under way, and works out the client.
// under way, works out the client, and starts the log line with what is
// known of the request.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
h.metrics.RequestStarted()
start := time.Now()
peer := peerAddress(r)
trusted := h.config.TrustedProxies
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted)
peerTrusted := isInside(peer, trusted)
forwardedFor := r.Header.Values("X-Forwarded-For")
client := clientAddress(peer, forwardedFor, trusted)
rq := &request{
h: h,
@@ -84,22 +97,37 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
out: &responseWriter{ResponseWriter: w},
client: client,
peer: peer,
peerTrusted: isInside(peer, trusted),
peerTrusted: peerTrusted,
start: start,
line: requestlog.Line{
Time: requestlog.FormatTime(start),
ClientIP: client.String(),
PeerIP: peer.String(),
Method: r.Method,
Host: r.Host,
Path: r.URL.EscapedPath(),
Query: r.URL.RawQuery,
Protocol: r.Proto,
Referer: r.Referer(),
UserAgent: r.UserAgent(),
Action: requestlog.ActionForward,
Time: requestlog.FormatTime(start),
Instance: h.config.InstanceName,
ClientIP: client.String(),
Method: r.Method,
Scheme: scheme(r, peerTrusted),
Host: r.Host,
Path: r.URL.EscapedPath(),
Query: r.URL.RawQuery,
Protocol: r.Proto,
Referer: r.Referer(),
UserAgent: r.UserAgent(),
RequestID: requestID(r, peerTrusted),
PeerIP: peer.String(),
ForwardedFor: strings.Join(forwardedFor, ", "),
ClientGroup: clientGroup(client).String(),
ContentType: r.Header.Get("Content-Type"),
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
HasAuthorization: len(r.Header.Values("Authorization")) > 0,
HasCookie: len(r.Header.Values("Cookie")) > 0,
Action: requestlog.ActionForward,
},
}
// A length of -1 is a body whose length was not announced.
if r.ContentLength > 0 {
rq.line.ContentLength = r.ContentLength
}
if r.Body != http.NoBody {
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
}
@@ -107,23 +135,47 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
return rq
}
// requestHeaders returns the headers of r that names lists, by name in
// lower case, each with its values joined by ", ". Authorization, Cookie
// and Set-Cookie are never among them, whatever names says.
func requestHeaders(r *http.Request, names []string) map[string]string {
headers := map[string]string{}
for _, name := range names {
switch name {
case "authorization", "cookie", "set-cookie":
continue
}
values := r.Header.Values(name)
if len(values) > 0 {
headers[name] = strings.Join(values, ", ")
}
}
return headers
}
// check is the one place where a request can be refused once its client
// is known, before its body is read or anything reaches the app. It
// returns nil to let the request through. The checks of checkClient come
// first, answered with SWWAF_BAN_RESPONSE, and then the size limit, so
// that a request the rate limits count is counted even when it is
// refused for its size. In observe mode a request checkClient refuses
// goes on to the size limit like any other. ctx is the request's own
// context.
// first, answered with SWWAF_BAN_RESPONSE, or 403 for a block rule, and
// then the size limit, so that a request the rate limits count is counted
// even when it is refused for its size. In observe mode a request
// checkClient refuses goes on to the size limit like any other. ctx is
// the request's own context.
func (rq *request) check(ctx context.Context) *refusal {
action := rq.checkClient(ctx)
if action != "" {
if !rq.h.config.Observe {
return rq.banResponse(action)
}
switch {
case action == "":
case rq.h.config.Observe:
// The log line names what enforce mode would have done.
rq.line.WouldAction = action
case action == requestlog.ActionRuleBlocked:
return &refusal{status: http.StatusForbidden, action: action}
default:
return rq.banResponse(action)
}
maxBytes := rq.h.config.RequestMaxBytes
@@ -145,8 +197,9 @@ func (rq *request) check(ctx context.Context) *refusal {
// client either refuses is not looked up, and then the country lists; a
// request any of them refuses is not counted for the rate limits. Then
// come the rate limits, unless the client is in
// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted.
// ctx is the request's own context.
// SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under
// SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted,
// and last the rule files. ctx is the request's own context.
func (rq *request) checkClient(ctx context.Context) string {
cfg := rq.h.config
if isInside(rq.client, cfg.AllowNets) {
@@ -167,11 +220,37 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionCountryDenied
}
if !isInside(rq.client, cfg.RateLimitExemptNets) && rq.limitBroken(now) {
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if !exempt && rq.limitBroken(now) {
return requestlog.ActionRateLimited
}
return ""
return rq.checkRules(now)
}
// pathExempt reports whether the rate limits leave out a request for u
// because of SWWAF_RATE_LIMIT_EXEMPT_PATHS: whether its path as sent, the
// path the app receives, not percent-decoded, starts with one of
// prefixes, so that /%61ssets/x is not under /assets/ for an app whose
// router matches the path as received. A request whose decoded path
// contains .. anywhere or a backslash, or whose path as sent holds an
// encoded slash (%2F or %2f), never is, since an app may act on it as a
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
// as one path segment, as Go's router does.
func pathExempt(u *url.URL, prefixes []string) bool {
decoded := u.Path
// EscapedPath is the path as the app receives it, not decoded.
sent := u.EscapedPath()
if strings.Contains(decoded, "..") || strings.Contains(decoded, `\`) ||
strings.Contains(strings.ToLower(sent), "%2f") {
return false
}
return slices.ContainsFunc(prefixes, func(prefix string) bool {
return strings.HasPrefix(sent, prefix)
})
}
// forward passes the request to the app and the app's answer back. ctx
@@ -182,7 +261,9 @@ func (rq *request) forward(ctx context.Context) {
rq.cancel = cancel
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
WroteRequest: rq.wroteRequest,
GotConn: rq.gotConn,
WroteRequest: rq.wroteRequest,
GotFirstResponseByte: rq.gotFirstResponseByte,
})
out := rq.in.WithContext(ctx)
@@ -205,7 +286,8 @@ func (rq *request) forward(ctx context.Context) {
}
// rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set.
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set.
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme
@@ -214,6 +296,7 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
// the query as the client sent it.
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
}
// modifyResponse looks at the app's answer before ReverseProxy passes it
@@ -227,6 +310,7 @@ func (rq *request) modifyResponse(res *http.Response) error {
// connection it takes over, not through rq.out.
rq.stopTimers()
rq.out.status = res.StatusCode
rq.line.Websocket = true
return nil
}
@@ -303,10 +387,14 @@ func (rq *request) answer(r refusal) {
}
// refuse records r, unless an earlier refusal was, and ends the request
// to the app.
// to the app, if one was made: smallwebwaf reads the body of a request
// it answers itself too.
func (rq *request) refuse(r refusal) {
rq.refused.CompareAndSwap(nil, &r)
rq.cancel()
if rq.cancel != nil {
rq.cancel()
}
}
// finish ends the request's timeouts, counts it in the metrics and writes
@@ -322,6 +410,10 @@ func (rq *request) finish() {
line := &rq.line
line.Status = rq.out.status
line.ResponseBytes = rq.out.bytes
header := rq.out.Header()
line.ResponseContentType = header.Get("Content-Type")
line.CacheControl = header.Get("Cache-Control")
line.Location = header.Get("Location")
if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load()
@@ -346,12 +438,18 @@ func (rq *request) finish() {
now := time.Now()
duration := now.Sub(rq.start)
line.DurationTotal = requestlog.Milliseconds(duration)
line.DurationChecks = timing(rq.start, rq.checked)
var upstreamDuration time.Duration
if !rq.upstreamStart.IsZero() {
upstreamDuration = now.Sub(rq.upstreamStart)
line.DurationUpstreamTotal = requestlog.Milliseconds(upstreamDuration)
line.DurationUpstreamTotal = new(requestlog.Milliseconds(upstreamDuration))
rq.mu.Lock()
line.DurationUpstreamConnect = timing(rq.upstreamStart, rq.connected)
line.DurationUpstreamFirstByte = timing(rq.upstreamStart, rq.answerStarted)
rq.mu.Unlock()
}
// Counted before the log line is written, so that the metrics count
@@ -364,6 +462,17 @@ func (rq *request) finish() {
}
}
// timing is the time from start to end in milliseconds, for one of the
// log line's timings, or nil when end is zero: what it times never
// happened.
func timing(start, end time.Time) *float64 {
if end.IsZero() {
return nil
}
return new(requestlog.Milliseconds(end.Sub(start)))
}
// addToHistory adds the request, which has ended, to its client's
// history.
func (rq *request) addToHistory() {
@@ -473,6 +582,24 @@ func (rq *request) bodyReceived() {
stopTimer(rq.clientRequestTimer)
}
// gotConn is called once there is a connection to the app, a new one or
// one kept open from an earlier request.
func (rq *request) gotConn(httptrace.GotConnInfo) {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.connected = time.Now()
}
// gotFirstResponseByte is called once the first byte of the app's answer
// has arrived.
func (rq *request) gotFirstResponseByte() {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.answerStarted = time.Now()
}
// wroteRequest is called once the app has been sent the whole request:
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
+368
View File
@@ -0,0 +1,368 @@
package proxy_test
import (
"io"
"maps"
"math"
"net/http"
"reflect"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// requestIDHeader carries the request's id.
requestIDHeader = "X-Request-ID"
// instance is the SWWAF_INSTANCE_NAME a test sets.
instance = "fsn1app1/gitea"
// ipv6Client is a client on IPv6, and ipv6Group the netblock the rate
// limits count it as.
ipv6Client = "2001:db8::7"
ipv6Group = "2001:db8::/64"
)
func TestLogLineHasEachFieldWhereItApplies(t *testing.T) {
t.Parallel()
received := make(chan string, 2) // the request ids the app received
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path != "/full" {
w.WriteHeader(http.StatusNoContent)
return
}
w.Header().Set("Content-Type", "text/html")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Location", "/elsewhere")
w.WriteHeader(http.StatusFound)
_, _ = io.WriteString(w, "moved")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
rateLimitExemptNets: localhost,
instanceName: instance,
logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE",
})
// This request comes from ipv6Client through a trusted proxy, with a
// body and each header the log line looks at, and is answered with a
// redirect.
conn := dial(t, addr)
send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+
forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+
forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+
"Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+
"Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+
"Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+
"Set-Cookie: secret-set-cookie\r\n\r\na=b")
wantStatus(t, readResponse(t, conn), http.StatusFound)
// A request's log line can come after its answer: each is waited for
// before the next request, so that the lines are in order.
full := out.requestLines(t, 1)[0]
// This one comes from 127.0.0.1, which the rate limits do not count,
// with a body of 4 bytes whose length it does not announce, so that its
// request_bytes is not its content_length, and no header the log line
// looks at, and is answered with 204 and no header.
conn = dial(t, addr)
send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+
"Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusNoContent)
bare := out.requestLines(t, 2)[1]
wantFullLine(t, full)
wantBareLine(t, bare)
for _, line := range []logLine{full, bare} {
got := <-received
if got != line.RequestID {
t.Errorf("the app received request id %q, the log line has %q",
got, line.RequestID)
}
}
if strings.Contains(out.text(), "secret") {
t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s",
out.text())
}
}
// wantFullLine checks the log line of the request with every header the
// line looks at. Its timings are checked by TestTimingsAreInOrder.
func wantFullLine(t *testing.T, line logLine) {
t.Helper()
headers := map[string]string{"accept": "text/html", "x-custom": "one, two"}
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure,
Host: appHost, Path: "/full", Protocol: protocol,
Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5,
RequestID: "from-traefik", PeerIP: localhost,
ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group,
ContentType: "application/x-www-form-urlencoded", ContentLength: 3,
RequestHeaders: headers, HasAuthorization: true, HasCookie: true,
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
}
// wantBareLine checks the log line of the request with none of them, and
// that the fields that do not apply to it are left out.
func wantBareLine(t *testing.T, line logLine) {
t.Helper()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: localhost, Method: http.MethodPost, Scheme: plain,
Host: appHost, Path: "/bare", Protocol: protocol,
Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID,
PeerIP: localhost, ClientGroup: localhost + "/32",
UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward,
})
if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" {
t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want)
}
for _, name := range []string{
"forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"cache_control", "location", "counts",
} {
_, present := line.fields[name]
if present {
t.Errorf("log line has %s, which does not apply", name)
}
}
}
// withTimings returns want with the timings of line.
func withTimings(line logLine, want requestlog.Line) requestlog.Line {
want.DurationTotal = line.DurationTotal
want.DurationChecks = line.DurationChecks
want.DurationUpstreamConnect = line.DurationUpstreamConnect
want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte
want.DurationUpstreamTotal = line.DurationUpstreamTotal
return want
}
func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) {
t.Parallel()
const hasAuthorization, hasCookie = "has_authorization", "has_cookie"
for _, tc := range []struct{ header, field, other string }{
{"Authorization", hasAuthorization, hasCookie},
{"Cookie", hasCookie, hasAuthorization},
} {
t.Run("only "+tc.header, func(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out := startProxy(t, app.URL, nil)
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(tc.header, "secret")
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLine(t)
_, otherPresent := line.fields[tc.other]
if line.fields[tc.field] != true || otherPresent {
t.Errorf("log line has %s %v and %s %v, want true and none",
tc.field, line.fields[tc.field], tc.other, line.fields[tc.other])
}
})
}
}
func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) {
t.Parallel()
const sentID = "from-traefik"
sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}}
trusted := map[string]string{trustedProxies: trustLocalhost}
for _, tc := range []struct {
name string
env map[string]string
header http.Header
// wantID is the request id logged, "" for a new one.
wantID, wantScheme string
}{
{"a trusted proxy's are kept", trusted, sent, sentID, secure},
{"without them, the id is new and the scheme http", trusted, nil, "", plain},
{"another peer's are replaced", nil, sent, "", plain},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
received := make(chan string, 2)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
})
addr, out := startProxy(t, app.URL, tc.env)
// Two requests, so that two new ids can be told apart.
ids := make([]string, 0, 2)
for i := range 2 {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
maps.Copy(req.Header, tc.header)
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLines(t, i+1)[i]
ids = append(ids, line.RequestID)
got := <-received
if line.RequestID != got || line.Scheme != tc.wantScheme {
t.Errorf("log line has request_id %q and scheme %q, and the "+
"app received id %q; want the same id and scheme %q",
line.RequestID, line.Scheme, got, tc.wantScheme)
}
}
switch {
case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID):
t.Errorf("request ids %q, want %q", ids, tc.wantID)
case tc.wantID == "" && (slices.Contains(ids, sentID) ||
slices.Contains(ids, "") || ids[0] == ids[1]):
t.Errorf("request ids %q, want two new ones", ids)
}
})
}
}
func TestTimingsAreInOrder(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
// The pauses set the times apart; a hold-up of the test only
// lengthens them.
time.Sleep(time.Millisecond)
w.WriteHeader(http.StatusOK)
_ = http.NewResponseController(w).Flush()
time.Sleep(time.Millisecond)
_, _ = io.WriteString(w, "done")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
denyNets: denied,
})
// Each log line is waited for before the next request, so that the
// lines are in order.
wantStatus(t, get(t, addr, "/"), http.StatusOK)
forwarded := out.requestLines(t, 1)[0]
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, denied)
wantStatus(t, do(t, req), http.StatusForbidden)
refused := out.requestLines(t, 2)[1]
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
health := out.requestLines(t, 3)[2]
// A request passed to the app has every timing; one refused, none of
// the app's; the health check, which runs no check, only the total.
wantTimings(t, forwarded, "duration_total", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total")
wantTimings(t, refused, "duration_total", "duration_checks")
wantTimings(t, health, "duration_total")
if t.Failed() {
return
}
// In whole microseconds, as they are logged, so that the sum below is
// exact.
total := microseconds(forwarded.DurationTotal)
checks := microseconds(*forwarded.DurationChecks)
connect := microseconds(*forwarded.DurationUpstreamConnect)
firstByte := microseconds(*forwarded.DurationUpstreamFirstByte)
upstream := microseconds(*forwarded.DurationUpstreamTotal)
// The checks end before the request is handed to the app, and the
// connection comes before the answer, which the app ends after a
// pause.
if checks+upstream > total || connect >= firstByte || firstByte >= upstream {
t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+
"first byte %d, upstream total %d", total, checks, connect, firstByte,
upstream)
}
if *refused.DurationChecks > refused.DurationTotal {
t.Errorf("refused request's checks took %v of %v milliseconds",
*refused.DurationChecks, refused.DurationTotal)
}
}
// wantTimings checks that the timings named are the only ones line has.
func wantTimings(t *testing.T, line logLine, want ...string) {
t.Helper()
var got []string
for name := range line.fields {
if strings.HasPrefix(name, "duration_") {
got = append(got, name)
}
}
slices.Sort(got)
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want)
}
}
// microseconds is a timing in whole microseconds.
func microseconds(milliseconds float64) int64 {
return int64(math.Round(milliseconds * 1000))
}
func TestLogsAnUpgradedConnection(t *testing.T) {
t.Parallel()
app := startApp(t, echoAfterUpgrade)
addr, out := startProxy(t, app.URL, nil)
conn := dial(t, addr)
send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols)
_ = conn.Close()
line := out.requestLine(t)
if line.fields["websocket"] != true {
t.Errorf("log line has websocket %v, want true", line.fields["websocket"])
}
}
+233
View File
@@ -0,0 +1,233 @@
package proxy_test
import (
"net/http"
"net/netip"
"os"
"path/filepath"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// testRules are the rules most tests here load: a block rule for
// /blocked and a ban rule for /.env.
const testRules = `
blocked path block ^/blocked$
probe path ban ^/\.env$
`
func TestEachRuleAction(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, "noted path log ^/\n"+testRules),
banResponse: "429",
})
start := clk.Now()
// A log rule notes its match, and lets the request through.
line := s.get(client, http.StatusOK, requestlog.ActionForward)
wantRuleIDs(t, line, "noted")
// A block rule refuses with 403, whatever SWWAF_BAN_RESPONSE is, and
// bans no one.
line = s.request(client, "/blocked", http.StatusForbidden,
requestlog.ActionRuleBlocked)
wantRuleIDs(t, line, "noted", "blocked")
s.get(client, http.StatusOK, requestlog.ActionForward)
// A ban rule refuses with SWWAF_BAN_RESPONSE, and bans the client for
// seven days, the default.
line = s.request(client, "/.env", http.StatusTooManyRequests, requestlog.ActionBanned)
wantRuleIDs(t, line, "noted", "probe")
if line.BanExpires != requestlog.FormatTime(start.Add(7*24*time.Hour)) {
t.Errorf("log line has ban_expires %q, want seven days on", line.BanExpires)
}
netblock := netip.MustParsePrefix(client + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: start.Add(7 * 24 * time.Hour),
Cause: bans.CauseAttack,
Reason: "matched the rule probe",
Notes: bans.Notes{
RuleID: "probe",
Target: "path",
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: "/.env",
Status: http.StatusTooManyRequests,
UserAgent: userAgent,
},
// The four requests up to and including the probe.
Requests: 4,
},
}
got := server.Ledger.Bans(netblock)
if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
// The next request is refused under the ban, without being checked
// against the rules, and makes the ban permanent.
clk.advance(time.Hour)
line = s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
wantRuleIDs(t, line)
if line.BanExpires != permanent {
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
}
clk.advance(365 * 24 * time.Hour)
s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
}
func TestNextClearSignOfAttackAfterABanBansPermanently(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
attackBanDuration: "1h",
})
// The first probe bans for SWWAF_ATTACK_BAN_DURATION.
line := s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != requestlog.FormatTime(clk.Now().Add(time.Hour)) {
t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires)
}
// Once that ban has run out without a request, the client is served,
// and its next probe bans it for good.
clk.advance(time.Hour)
s.get(client, http.StatusOK, requestlog.ActionForward)
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != permanent {
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
}
}
func TestRulesComeAfterTheOtherChecks(t *testing.T) {
t.Parallel()
const (
allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
exempt = "192.0.2.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
)
s, _, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
allowNets: allowed,
rateLimitExemptNets: exempt,
rateLimitPerMinute: "1",
})
// A client in SWWAF_ALLOW_NETS is not checked.
line := s.request(allowed, "/.env", http.StatusOK, requestlog.ActionForward)
wantRuleIDs(t, line)
// A probe over the rate limit breaks the limit before any rule sees
// it.
s.get(client, http.StatusOK, requestlog.ActionForward)
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionRateLimited)
wantRuleIDs(t, line)
limitBan := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(limitBan) != 1 || limitBan[0].Cause != bans.CauseLimit {
t.Errorf("bans %+v, want one for a broken limit", limitBan)
}
// A client the rate limits do not apply to is still checked.
s.get(exempt, http.StatusOK, requestlog.ActionForward)
s.request(exempt, "/.env", http.StatusForbidden, requestlog.ActionBanned)
}
func TestObserveModeLogsWhatTheRulesWouldDo(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
mode: observe,
})
line := s.request(client, "/blocked", http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionRuleBlocked)
wantRuleIDs(t, line, "blocked")
line = s.request(client, "/.env", http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionBanned)
wantRuleIDs(t, line, "probe")
if line.BanExpires != "" {
t.Errorf("log line has ban_expires %q, want none", line.BanExpires)
}
// No ban was made.
line = s.get(client, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, "")
if got := server.Ledger.Snapshot(); len(got) != 0 {
t.Errorf("bans %+v, want none", got)
}
}
func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
t.Parallel()
const scraper = "192.0.2.200"
s, _, _ := startWithClock(t, "", map[string]string{
rulesDir: writeRules(t, testRules),
metricsToken: token,
})
s.request(client, "/blocked", http.StatusForbidden, requestlog.ActionRuleBlocked)
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
wantMetric(t, metrics,
`smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
}
// writeRules writes content as a rule file into a new directory, and
// returns the directory.
func writeRules(t *testing.T, content string) string {
t.Helper()
dir := t.TempDir()
err := os.WriteFile(filepath.Join(dir, "test.rules"), []byte(content), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
return dir
}
// wantRuleIDs checks the request log line's rule_ids.
func wantRuleIDs(t *testing.T, line logLine, want ...string) {
t.Helper()
if !slices.Equal(line.RuleIDs, want) {
t.Errorf("log line has rule_ids %v, want %v", line.RuleIDs, want)
}
}
+39
View File
@@ -0,0 +1,39 @@
package proxy
import (
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
// checkRules checks the request against the rules of the rule files at
// now, notes the ids of those it matches in the log line, and returns the
// action of the rule that refuses it, ActionRuleBlocked for a block rule
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans
// the client's netblock for a clear sign of attack, or in observe mode
// raises the alert for the ban it would have made.
func (rq *request) checkRules(now time.Time) string {
matched := rq.h.rules.Match(rq.in)
for _, rule := range matched {
rq.line.RuleIDs = append(rq.line.RuleIDs, rule.ID)
rq.h.metrics.RuleMatched(rule.ID, rule.Action)
}
if len(matched) == 0 {
return ""
}
// Only the last rule matched can refuse the request.
switch last := matched[len(matched)-1]; last.Action {
case rules.ActionBlock:
return requestlog.ActionRuleBlocked
case rules.ActionBan:
rq.banForAttack(now, last)
return requestlog.ActionBanned
default:
return ""
}
}
+36 -9
View File
@@ -149,26 +149,39 @@ type Hit struct {
Requests float64
}
// Counts are a client's requests in the minute, the hour and the day that
// end at a request, that request included.
type Counts struct {
Minute float64 `json:"minute"`
Hour float64 `json:"hour"`
Day float64 `json:"day"`
}
// Count counts a request from client at now, in every window, whether or
// not it is refused. It reports whether the request takes the client over
// a limit, and the window whose limit it goes over, the shortest if it is
// over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
// not it is refused, and returns the client's requests in each window. It
// reports whether the request takes the client over a limit, and the
// window whose limit it goes over, the shortest if it is over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
var hit Hit
var (
requests [3]float64
hit Hit
)
for i, b := range l.get(client).buckets() {
w := l.windows[i]
requests := b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
requests[i] = b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
}
}
return hit, hit.Window != ""
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
return counts, hit, hit.Window != ""
}
// Reset sets client's counts in every window back to zero. Its history
@@ -242,6 +255,20 @@ func (l *Limiter) Requests(netblock netip.Prefix) int64 {
return requests
}
// Client returns client as the table holds it, and whether it does. It is
// not a request from client, and leaves when it was last seen unchanged.
func (l *Limiter) Client(client netip.Prefix) (Client, bool) {
l.mu.Lock()
defer l.mu.Unlock()
c, seen := l.clients.Peek(client)
if !seen {
return Client{}, false
}
return *c, true
}
// Len returns how many clients are in the table.
func (l *Limiter) Len() int {
l.mu.Lock()
+26 -3
View File
@@ -62,14 +62,14 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight()
for range limit {
_, over := limiter.Count(client, start)
_, _, over := limiter.Count(client, start)
if over {
t.Fatal("a request within the limit is over it")
}
}
// Over both limits; the minute's is named, with the four requests.
hit, over := limiter.Count(client, start)
_, hit, over := limiter.Count(client, start)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want {
@@ -78,6 +78,29 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
}
}
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range 3 {
limiter.Count(client, start)
}
// A quarter into the next hour, the minute has only this request. The
// hour still covers three quarters of the bucket before, with its three
// requests, which count 2.25, and this one: 3.25. The day covers all
// four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want {
t.Errorf("counts %+v, want %+v", counts, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel()
@@ -238,7 +261,7 @@ func wantCount(
) {
t.Helper()
hit, _ := limiter.Count(client, now)
_, hit, _ := limiter.Count(client, now)
if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want)
+310
View File
@@ -0,0 +1,310 @@
// Package remotelog sends the lines smallwebwaf writes on stdout to the
// remote log endpoint, SWWAF_LOG_REMOTE_URL, as the "Request log" section
// of SPEC.md describes: each line as the message of an RFC 5424 syslog
// record, over UDP, TCP or TLS. Lines wait in a bounded buffer, so a slow
// or unreachable endpoint never holds up a request or stdout.
package remotelog
import (
"bytes"
"context"
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"log/slog"
"net"
"net/url"
"os"
"strconv"
"sync/atomic"
"syscall"
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The forms of SWWAF_LOG_REMOTE_URL, by its scheme.
const (
SchemeUDP = "syslog+udp"
SchemeTCP = "syslog+tcp"
SchemeTLS = "syslog+tls"
)
// A record's priority is the number of its facility times the number of
// severities there are, plus the number of its severity. Every record's
// severity is informational.
const (
severities = 8
informational = 6
)
const (
// dialTimeout bounds connecting to the endpoint, the TLS handshake
// included.
dialTimeout = 10 * time.Second
// After a failed attempt to connect, or a connection on which a record
// fails, the next attempt to connect is made a second later, and
// retryDelayFactor times as long after each further failure in a row,
// up to a minute. A connection that fails after it has stayed up for
// resetRetryDelayAfter ends the row.
firstRetryDelay = time.Second
retryDelayFactor = 2
maxRetryDelay = time.Minute
resetRetryDelayAfter = time.Minute
)
// Params are what New needs.
type Params struct {
// URL is the endpoint (SWWAF_LOG_REMOTE_URL): SchemeUDP, SchemeTCP or
// SchemeTLS, a host and a port.
URL *url.URL
// RootCAs are the certificates a SchemeTLS endpoint's certificate
// must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil for the host's.
RootCAs *x509.CertPool
// Buffer is the most lines held while they wait to be sent
// (SWWAF_LOG_REMOTE_BUFFER).
Buffer int
// Facility is the number of the records' syslog facility
// (SWWAF_LOG_REMOTE_FACILITY), and AppName their APP-NAME
// (SWWAF_LOG_REMOTE_APP_NAME).
Facility int
AppName string
}
// Sender sends lines to the endpoint. Write puts them in its buffer, and
// Run sends them from there.
type Sender struct {
url *url.URL
tlsConfig *tls.Config
// beforeTime and afterTime are the parts of every record's header
// before and after its time, as RFC 5424 lays the header out.
beforeTime string
afterTime string
// records is the buffer: each line's record, framed to be sent.
records chan []byte
sent atomic.Int64
dropped atomic.Int64
}
// New returns a Sender for the endpoint params.URL.
func New(params Params) *Sender {
hostname, err := os.Hostname()
if err != nil || hostname == "" {
hostname = "-" // RFC 5424's value for a field that has none
}
priority := params.Facility*severities + informational
return &Sender{
url: params.URL,
tlsConfig: &tls.Config{
RootCAs: params.RootCAs,
MinVersion: tls.VersionTLS12,
},
// The 1 is the version of the format. The process id, the message
// id and the structured data have no value.
beforeTime: "<" + strconv.Itoa(priority) + ">1 ",
afterTime: " " + hostname + " " + params.AppName + " - - - ",
records: make(chan []byte, params.Buffer),
}
}
// Write puts each line in p in the buffer, as the message of a record of
// its own, and never waits: when the buffer is full, the oldest record in
// it is dropped to make room. It is safe for concurrent use.
func (s *Sender) Write(p []byte) (int, error) {
at := requestlog.FormatTime(time.Now())
for line := range bytes.Lines(p) {
line = bytes.TrimSuffix(line, []byte("\n"))
if len(line) > 0 {
s.put(s.record(at, line))
}
}
return len(p), nil
}
// Sent is how many records have been sent.
func (s *Sender) Sent() int64 {
return s.sent.Load()
}
// Dropped is how many records were dropped: the oldest in a full buffer,
// and those whose sending failed.
func (s *Sender) Dropped() int64 {
return s.dropped.Load()
}
// Depth is how many records are in the buffer.
func (s *Sender) Depth() int {
return len(s.records)
}
// Run connects to the endpoint and sends each record as it comes into the
// buffer, until ctx is done. Then it sends the records still in the buffer,
// on the connection open at that time or, if there is none, on a new one,
// until none is left or one fails, and returns. How long it may take over
// that is for the caller to bound.
//
// A connection on which a record fails is closed and the record dropped.
// That failure, like a failed attempt to connect, is logged to processLog
// and followed by the next attempt after firstRetryDelay, retryDelayFactor
// times as long after each further failure in a row up to maxRetryDelay,
// and firstRetryDelay again after a connection that stayed up for
// resetRetryDelayAfter. Meanwhile the records wait in the buffer.
func (s *Sender) Run(ctx context.Context, processLog *slog.Logger) {
conn := s.send(ctx, processLog)
if conn == nil && len(s.records) > 0 {
conn, _ = s.dial(context.WithoutCancel(ctx))
}
if conn == nil {
return
}
defer func() {
_ = conn.Close()
}()
for {
select {
case record := <-s.records:
if s.write(conn, record) != nil {
return
}
default:
return
}
}
}
// record returns line as an RFC 5424 record made at the time at, framed
// for the endpoint: on its own over UDP, since each datagram holds one,
// and over TCP and TLS after its length in bytes and a space, the
// octet-counted framing of RFC 6587 and RFC 5425.
func (s *Sender) record(at string, line []byte) []byte {
record := make([]byte, 0, len(s.beforeTime)+len(at)+len(s.afterTime)+len(line))
record = append(record, s.beforeTime...)
record = append(record, at...)
record = append(record, s.afterTime...)
record = append(record, line...)
if s.url.Scheme == SchemeUDP {
return record
}
return append([]byte(strconv.Itoa(len(record))+" "), record...)
}
// put adds record to the buffer, first dropping the oldest record in it
// while it is full.
func (s *Sender) put(record []byte) {
for {
select {
case s.records <- record:
return
default:
}
select {
case <-s.records:
s.dropped.Add(1)
default:
}
}
}
// send connects to the endpoint and sends each record as it comes into
// the buffer, until ctx is done, and returns the connection then open, or
// nil.
func (s *Sender) send(ctx context.Context, processLog *slog.Logger) net.Conn {
delay := firstRetryDelay
for {
conn, err := s.dial(ctx)
if ctx.Err() != nil {
return conn
}
if err == nil {
connected := time.Now()
err = s.sendOn(ctx, conn)
if err == nil {
return conn
}
_ = conn.Close()
if time.Since(connected) >= resetRetryDelayAfter {
delay = firstRetryDelay
}
}
processLog.Warn("sending to SWWAF_LOG_REMOTE_URL failed",
"error", err.Error(), "connecting_again_in", delay.String())
select {
case <-time.After(delay):
case <-ctx.Done():
return nil
}
delay = min(retryDelayFactor*delay, maxRetryDelay)
}
}
// sendOn sends each record on conn as it comes into the buffer, until one
// fails, whose error it returns, or ctx is done.
func (s *Sender) sendOn(ctx context.Context, conn net.Conn) error {
for {
select {
case record := <-s.records:
err := s.write(conn, record)
if err != nil {
return err
}
case <-ctx.Done():
return nil
}
}
}
// write sends record on conn, and counts it as sent or, if that fails,
// as dropped. A record too long for one UDP datagram is dropped without
// an error, since the connection has not failed: a long request must not
// hold up the lines after it.
func (s *Sender) write(conn net.Conn, record []byte) error {
_, err := conn.Write(record)
if err != nil {
s.dropped.Add(1)
if errors.Is(err, syscall.EMSGSIZE) {
return nil
}
return fmt.Errorf("send a record: %w", err)
}
s.sent.Add(1)
return nil
}
// dial connects to the endpoint.
func (s *Sender) dial(ctx context.Context) (net.Conn, error) {
dialer := &net.Dialer{Timeout: dialTimeout}
switch s.url.Scheme {
case SchemeUDP:
return dialer.DialContext(ctx, "udp", s.url.Host)
case SchemeTLS:
tlsDialer := &tls.Dialer{NetDialer: dialer, Config: s.tlsConfig}
return tlsDialer.DialContext(ctx, "tcp", s.url.Host)
default:
return dialer.DialContext(ctx, "tcp", s.url.Host)
}
}
+640
View File
@@ -0,0 +1,640 @@
package remotelog_test
import (
"bufio"
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"fmt"
"io"
"log/slog"
"math/big"
"net"
"net/url"
"os"
"slices"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// The tests run in a synctest bubble, where the time package runs on a
// clock of the test's own, which starts at 2000-01-01T00:00:00Z: a wait
// lasts exactly as long as it should, however slowly the test process
// runs, and synctest.Wait returns once the sender has done all it can
// before time passes. The endpoint is a listener on the loopback address.
// A test reads from it only once the records are on their way, and checks
// the sender's counts first, since a goroutine of the bubble that waits on
// the network keeps that clock from moving on. For the same reason the
// endpoint that refuses connections, a tlsEndpoint, runs outside the
// bubble: a sender connecting over TLS waits on the endpoint's answer.
const (
// started is the time a record made as a test starts gives.
started = "2000-01-01T00:00:00.000Z"
appName = "fsn1app1/gitea"
// local0 is the number of the default facility, and local0Info the
// priority of its records.
local0 = 16
local0Info = "<134>"
// loopback is where the endpoints listen.
loopback = "127.0.0.1:0"
)
func TestRecordsOverUDPGoOnePerDatagram(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = endpoint.Close() })
sender, _, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
_, _ = sender.Write([]byte(`{"type":"request"}` + "\n" + `{"type":"process"}` + "\n"))
synctest.Wait()
wantCounts(t, sender, 2, 0, 0)
for _, line := range []string{`{"type":"request"}`, `{"type":"process"}`} {
datagram := make([]byte, 1024)
n, _, err := endpoint.ReadFrom(datagram)
if err != nil {
t.Fatalf("read: %v", err)
}
want := record(t, local0Info, appName, line)
if string(datagram[:n]) != want {
t.Errorf("datagram %q, want %q", datagram[:n], want)
}
}
})
}
func TestRecordsOverTCPAreOctetCountedWithTheirFacility(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
endpointParams := params(remotelog.SchemeTCP, endpoint.Addr())
endpointParams.Facility = 19 // local3
endpointParams.AppName = "gitea"
sender, _, _ := run(t, endpointParams)
_, _ = sender.Write([]byte("first\nsecond\n"))
synctest.Wait()
wantCounts(t, sender, 2, 0, 0)
frames := bufio.NewReader(accept(t, endpoint))
wantFrame(t, frames, record(t, "<158>", "gitea", "first"))
wantFrame(t, frames, record(t, "<158>", "gitea", "second"))
})
}
func TestStalledEndpointHoldsUpNoWriteAndOldestRecordsAreDropped(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
certificate, roots := testCertificate(t)
endpoint := listen(t)
endpointParams := params(remotelog.SchemeTLS, endpoint.Addr())
endpointParams.RootCAs = roots
endpointParams.Buffer = 3
sender, _, _ := run(t, endpointParams)
// The sender connects, and its TLS handshake waits for an answer
// the endpoint does not give yet.
conn := accept(t, endpoint)
var stdout bytes.Buffer
out := io.MultiWriter(&stdout, sender)
for i := range 5 {
_, _ = fmt.Fprintf(out, "line %d\n", i+1)
}
if stdout.String() != "line 1\nline 2\nline 3\nline 4\nline 5\n" {
t.Errorf("stdout has %q", stdout.String())
}
wantCounts(t, sender, 0, 2, 3)
// Once the endpoint answers, the three newest records are sent.
server := tls.Server(conn, &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
})
err := server.HandshakeContext(t.Context())
if err != nil {
t.Fatalf("handshake: %v", err)
}
synctest.Wait()
wantCounts(t, sender, 3, 2, 0)
frames := bufio.NewReader(server)
for _, line := range []string{"line 3", "line 4", "line 5"} {
wantFrame(t, frames, record(t, local0Info, appName, line))
}
})
}
func TestReconnectsWithBackoffAfterTheEndpointGoesAway(t *testing.T) {
t.Parallel()
certificate, roots := testCertificate(t)
endpoint := startTLSEndpoint(t, certificate)
synctest.Test(t, func(t *testing.T) {
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
endpointParams.RootCAs = roots
sender, logged, _ := run(t, endpointParams)
_, _ = sender.Write([]byte("one\n"))
synctest.Wait()
wantCounts(t, sender, 1, 0, 0)
conn := endpoint.next(t)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "one"))
// The endpoint goes away: it closes the connection, and refuses the
// next ones. The sender notices when a record fails, and tries to
// connect again a second later, then two seconds after that.
endpoint.refusing.Store(true)
_ = conn.Close()
writeUntilDropped(t, sender, 1)
sent := sender.Sent()
_, _ = sender.Write([]byte("two\n"))
time.Sleep(time.Second)
synctest.Wait()
endpoint.refusing.Store(false)
time.Sleep(2*time.Second - time.Nanosecond)
synctest.Wait()
wantCounts(t, sender, sent, 1, 1)
// The endpoint is back, and the record waiting is sent.
time.Sleep(time.Nanosecond)
synctest.Wait()
wantCounts(t, sender, sent+1, 1, 0)
conn = endpoint.next(t)
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "two"))
wantRetries(t, logged, "1s", "2s")
})
}
func TestAConnectionClosedAtOnceIsMadeAgainAfterAGrowingDelay(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
// The endpoint closes each connection as soon as it takes it. The
// sender notices when a record fails, and connects again a second
// later, then two seconds after that, then four.
delays := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second}
for i, delay := range delays {
_ = accept(t, endpoint).Close()
writeUntilDropped(t, sender, int64(i+1))
wantConnectedAgainAfter(t, sender, delay)
}
wantRetries(t, logged, "1s", "2s", "4s")
})
}
func TestTheDelayStartsAgainAfterAConnectionThatStayedUpAMinute(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint := listen(t)
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
_ = accept(t, endpoint).Close()
writeUntilDropped(t, sender, 1)
wantConnectedAgainAfter(t, sender, time.Second)
// A connection that fails just short of a minute after it was made
// leaves the delay growing.
conn := accept(t, endpoint)
time.Sleep(time.Minute - time.Nanosecond)
_ = conn.Close()
writeUntilDropped(t, sender, 2)
wantConnectedAgainAfter(t, sender, 2*time.Second)
// One that fails a minute after it was made starts it again from a
// second.
conn = accept(t, endpoint)
time.Sleep(time.Minute)
_ = conn.Close()
writeUntilDropped(t, sender, 3)
wantConnectedAgainAfter(t, sender, time.Second)
wantRetries(t, logged, "1s", "2s", "1s")
})
}
func TestALineTooLongForADatagramIsDroppedAlone(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = endpoint.Close() })
sender, logged, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
// With its header, the first line's record is longer than the 65507
// bytes a UDP datagram over IPv4 holds. It is dropped, nothing is
// logged, and the next line is sent at once.
_, _ = sender.Write([]byte(strings.Repeat("x", 65507) + "\nnext\n"))
synctest.Wait()
wantCounts(t, sender, 1, 1, 0)
wantRetries(t, logged)
datagram := make([]byte, 1024)
n, _, err := endpoint.ReadFrom(datagram)
if err != nil {
t.Fatalf("read: %v", err)
}
want := record(t, local0Info, appName, "next")
if string(datagram[:n]) != want {
t.Errorf("datagram %q, want %q", datagram[:n], want)
}
})
}
func TestRecordsWaitingAtTheStopAreSent(t *testing.T) {
t.Parallel()
certificate, roots := testCertificate(t)
endpoint := startTLSEndpoint(t, certificate)
synctest.Test(t, func(t *testing.T) {
// The endpoint refuses the sender's first connection: it fails to
// connect, and waits a second to try again.
endpoint.refusing.Store(true)
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
endpointParams.RootCAs = roots
sender, logged, stop := run(t, endpointParams)
synctest.Wait()
wantRetries(t, logged, "1s")
_, _ = sender.Write([]byte("one\ntwo\n"))
endpoint.refusing.Store(false)
// Stopped before that second is over, it connects to send them.
stop()
wantCounts(t, sender, 2, 0, 0)
frames := bufio.NewReader(endpoint.next(t))
wantFrame(t, frames, record(t, local0Info, appName, "one"))
wantFrame(t, frames, record(t, local0Info, appName, "two"))
})
}
// output collects what the sender logs.
type output struct {
mu sync.Mutex
buf bytes.Buffer
}
// Write adds lines the sender logs.
func (o *output) Write(p []byte) (int, error) {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.Write(p)
}
// text returns everything logged so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// params returns the settings of a Sender for the endpoint at addr, in
// the form scheme names: room for ten lines, the default facility, and
// appName.
func params(scheme string, addr net.Addr) remotelog.Params {
return remotelog.Params{
URL: &url.URL{Scheme: scheme, Host: addr.String()},
Buffer: 10,
Facility: local0,
AppName: appName,
}
}
// run runs a Sender with settings until the test ends or the function
// it returns is called, which waits for Run to return. It returns the
// Sender, and what it logs.
func run(t *testing.T, settings remotelog.Params) (*remotelog.Sender, *output, func()) {
t.Helper()
sender := remotelog.New(settings)
logged := &output{}
ctx, cancel := context.WithCancel(t.Context())
ran := make(chan struct{})
go func() {
sender.Run(ctx, slog.New(slog.NewJSONHandler(logged, nil)))
close(ran)
}()
stop := func() {
cancel()
<-ran
}
t.Cleanup(stop)
return sender, logged, stop
}
// listen returns a TCP listener on the loopback address, closed when the
// test ends.
func listen(t *testing.T) net.Listener {
t.Helper()
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", loopback)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() { _ = listener.Close() })
return listener
}
// accept returns the next connection to listener, closed when the test
// ends.
func accept(t *testing.T, listener net.Listener) net.Conn {
t.Helper()
conn, err := listener.Accept()
if err != nil {
t.Fatalf("accept: %v", err)
}
t.Cleanup(func() { _ = conn.Close() })
return conn
}
// tlsEndpoint is a syslog+tls endpoint on the loopback address, which a
// test starts outside its bubble. It keeps its listener until the test
// ends, and either takes each connection or refuses it.
type tlsEndpoint struct {
addr net.Addr
// refusing is set while the endpoint closes each connection before the
// TLS handshake, which fails the sender's attempt to connect.
refusing atomic.Bool
// conns are the connections it has taken, after the handshake.
conns chan net.Conn
}
// startTLSEndpoint starts a tlsEndpoint with certificate, which takes
// connections until it is told to refuse them.
func startTLSEndpoint(t *testing.T, certificate tls.Certificate) *tlsEndpoint {
t.Helper()
listener := listen(t)
endpoint := &tlsEndpoint{addr: listener.Addr(), conns: make(chan net.Conn, 10)}
config := &tls.Config{
Certificates: []tls.Certificate{certificate},
MinVersion: tls.VersionTLS12,
}
go func() {
for {
conn, err := listener.Accept()
if err != nil {
return
}
server := tls.Server(conn, config)
if endpoint.refusing.Load() || server.HandshakeContext(t.Context()) != nil {
_ = conn.Close()
continue
}
endpoint.conns <- server
}
}()
return endpoint
}
// next returns the next connection the endpoint has taken, closed when
// the test ends.
func (e *tlsEndpoint) next(t *testing.T) net.Conn {
t.Helper()
conn := <-e.conns
t.Cleanup(func() { _ = conn.Close() })
return conn
}
// record returns the record of line made as the test started, with the
// priority and the app name given.
func record(t *testing.T, priority, app, line string) string {
t.Helper()
hostname, err := os.Hostname()
if err != nil || hostname == "" {
hostname = "-"
}
return priority + "1 " + started + " " + hostname + " " + app + " - - - " + line
}
// wantFrame reads the next octet-counted frame from frames, and checks
// that it holds want.
func wantFrame(t *testing.T, frames *bufio.Reader, want string) {
t.Helper()
count, err := frames.ReadString(' ')
if err != nil {
t.Fatalf("read a frame's length: %v", err)
}
length, err := strconv.Atoi(strings.TrimSuffix(count, " "))
if err != nil {
t.Fatalf("frame starts %q, not with its length", count)
}
got := make([]byte, length)
_, err = io.ReadFull(frames, got)
if err != nil {
t.Fatalf("read a frame: %v", err)
}
if string(got) != want {
t.Errorf("frame %q, want %q", got, want)
}
}
// wantCounts checks the records sender has sent, dropped and holds in
// its buffer.
func wantCounts(
t *testing.T, sender *remotelog.Sender, sent, dropped int64, depth int,
) {
t.Helper()
if sender.Sent() != sent || sender.Dropped() != dropped || sender.Depth() != depth {
t.Fatalf("sent %d, dropped %d, %d in the buffer; want %d, %d and %d",
sender.Sent(), sender.Dropped(), sender.Depth(), sent, dropped, depth)
}
}
// writeUntilDropped writes a line at a time until the count of records
// sender has dropped reaches dropped. The records it sends on a
// connection the endpoint has closed are lost before one fails; how many
// depends on when the endpoint's host answers that the connection is
// gone.
func writeUntilDropped(t *testing.T, sender *remotelog.Sender, dropped int64) {
t.Helper()
for sender.Dropped() < dropped {
_, _ = sender.Write([]byte("lost\n"))
synctest.Wait()
}
}
// wantConnectedAgainAfter writes a line while the sender waits to connect
// again, and checks that it connects, and takes the line from the buffer,
// only once delay is over.
func wantConnectedAgainAfter(
t *testing.T, sender *remotelog.Sender, delay time.Duration,
) {
t.Helper()
_, _ = sender.Write([]byte("waiting\n"))
time.Sleep(delay - time.Nanosecond)
synctest.Wait()
if sender.Depth() != 1 {
t.Fatalf("connected again before %v", delay)
}
time.Sleep(time.Nanosecond)
synctest.Wait()
if sender.Depth() != 0 {
t.Fatalf("not connected again after %v", delay)
}
}
// wantRetries checks that the sender logged a failure, of an attempt to
// connect or of a connection, for each of delays, the time until the next
// attempt, in order, and logged nothing else.
func wantRetries(t *testing.T, logged *output, delays ...string) {
t.Helper()
var got []string
for line := range strings.Lines(logged.text()) {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil || fields["msg"] != "sending to SWWAF_LOG_REMOTE_URL failed" {
t.Fatalf("logged %q", line)
}
delay, _ := fields["connecting_again_in"].(string)
got = append(got, delay)
}
if !slices.Equal(got, delays) {
t.Errorf("logged failures to connect again in %v, want %v", got, delays)
}
}
// testCertificate returns a certificate for 127.0.0.1 that is its own
// CA, and a pool that holds it. It is valid on the bubble's clock, which
// starts at 2000-01-01T00:00:00Z.
func testCertificate(t *testing.T) (tls.Certificate, *x509.CertPool) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("generate a key: %v", err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "smallwebwaf test CA"},
NotBefore: time.Date(1999, 12, 31, 0, 0, 0, 0, time.UTC),
NotAfter: time.Date(2000, 1, 2, 0, 0, 0, 0, time.UTC),
IsCA: true,
BasicConstraintsValid: true,
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)},
}
der, err := x509.CreateCertificate(rand.Reader, template, template,
&key.PublicKey, key)
if err != nil {
t.Fatalf("create a certificate: %v", err)
}
certificate, err := x509.ParseCertificate(der)
if err != nil {
t.Fatalf("parse the certificate: %v", err)
}
roots := x509.NewCertPool()
roots.AddCert(certificate)
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots
}
+81 -26
View File
@@ -9,6 +9,8 @@ import (
"io"
"log/slog"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// The action a request line names: what smallwebwaf did with the
@@ -26,8 +28,12 @@ const (
// ActionRateLimited is a request refused because it took its client
// over a rate limit, which bans the client.
ActionRateLimited = "rate_limited"
// ActionBanned is a request refused because a ban covers its client.
// ActionBanned is a request refused because a ban covers its client,
// or because it matched a ban rule, which bans the client.
ActionBanned = "banned"
// ActionRuleBlocked is a request refused because it matched a block
// rule.
ActionRuleBlocked = "rule_blocked"
// ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS.
ActionDenied = "denied"
@@ -45,32 +51,74 @@ const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
// Line is one request's line in the request log. The field names are
// those of the "Request log" section of SPEC.md.
// Line is one request's line in the request log. The field names, and
// their order, are those of the "Request log" section of SPEC.md. A field
// that may not apply to a request is left out of its line when it does
// not.
//
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Line struct {
Type string `json:"type"`
Time string `json:"time"`
ClientIP string `json:"client_ip"`
PeerIP string `json:"peer_ip"`
Country string `json:"country"`
Method string `json:"method"`
Host string `json:"host"`
Path string `json:"path"`
Query string `json:"query"`
Protocol string `json:"protocol"`
Status int `json:"status"`
UpstreamStatus int `json:"upstream_status,omitempty"`
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
Referer string `json:"referer"`
UserAgent string `json:"user_agent"`
Action string `json:"action"`
Type string `json:"type"`
// The standard web log fields. Scheme is how the client reached
// smallwebwaf, or the trusted proxy in front of it.
Time string `json:"time"`
Instance string `json:"instance"`
ClientIP string `json:"client_ip"`
Method string `json:"method"`
Scheme string `json:"scheme"`
Host string `json:"host"`
Path string `json:"path"`
Query string `json:"query"`
Protocol string `json:"protocol"`
Status int `json:"status"`
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
Referer string `json:"referer"`
UserAgent string `json:"user_agent"`
// Request detail. RequestID is the X-Request-ID a trusted proxy sent,
// or a new one, and is sent on to the app. ForwardedFor is the
// X-Forwarded-For header as received. ClientGroup is the netblock the
// client is counted as.
RequestID string `json:"request_id"`
PeerIP string `json:"peer_ip"`
ForwardedFor string `json:"forwarded_for,omitempty"`
ClientGroup string `json:"client_group"`
Country string `json:"country"`
ContentType string `json:"content_type,omitempty"`
// ContentLength is the length of its body the request announced.
ContentLength int64 `json:"content_length,omitempty"`
// RequestHeaders are the headers SWWAF_LOG_REQUEST_HEADERS names that
// the request carried, by name in lower case.
RequestHeaders map[string]string `json:"request_headers,omitempty"`
HasAuthorization bool `json:"has_authorization,omitempty"`
HasCookie bool `json:"has_cookie,omitempty"`
// Websocket is true when the connection was upgraded, as for a
// WebSocket.
Websocket bool `json:"websocket,omitempty"`
// Response detail, from the headers of the answer: the app's, as
// passed on, or those of smallwebwaf's own. Aborted is true when the
// client went away early.
ResponseContentType string `json:"response_content_type,omitempty"`
UpstreamStatus int `json:"upstream_status,omitempty"`
CacheControl string `json:"cache_control,omitempty"`
Location string `json:"location,omitempty"`
Aborted bool `json:"aborted,omitempty"`
// The decision.
Action string `json:"action"`
// WouldAction is, in observe mode, the action enforce mode would have
// taken with a request it would have refused: ActionDenied,
// ActionBanned, ActionCountryDenied or ActionRateLimited.
// ActionBanned, ActionCountryDenied, ActionRateLimited or
// ActionRuleBlocked.
WouldAction string `json:"would_action,omitempty"`
// Counts are the client's requests as the rate limits counted them
// with this one, for a request they counted.
Counts ratelimit.Counts `json:"counts,omitzero"`
// RuleIDs are the ids of the rule file rules the request matched.
RuleIDs []string `json:"rule_ids,omitempty"`
// LimitHit is the window whose rate limit the request went over:
// minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"`
@@ -79,11 +127,18 @@ type Line struct {
// BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent".
BanExpires string `json:"ban_expires,omitempty"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"`
// DurationTotal and DurationUpstreamTotal are in milliseconds.
DurationTotal float64 `json:"duration_total"`
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"`
// The timings, in milliseconds. DurationChecks is the time until the
// checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte
// and DurationUpstreamTotal run from when the request was handed to the
// app: until there was a connection to it, until the first byte of its
// answer arrived, and until the end. Each but DurationTotal is nil for
// a request that did not get that far.
DurationTotal float64 `json:"duration_total"`
DurationChecks *float64 `json:"duration_checks,omitempty"`
DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"`
DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"`
DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"`
}
// Write writes line to w as one JSON line marked "type":"request".
+5 -1
View File
@@ -50,7 +50,11 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
}
unset := []string{
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted",
"forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"upstream_status", "cache_control", "location", "aborted", "counts",
"limit_hit", "offence", "ban_expires", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total",
}
for _, name := range unset {
+483
View File
@@ -0,0 +1,483 @@
// Package rules reads the rule files: the plain text files in
// SWWAF_RULES_DIR, one rule to a line, that each request is checked
// against, as the "Rule files" section of SPEC.md describes. They are read
// at start, and again once the directory has had no change for a short
// time after one is edited, added or removed.
package rules
import (
"context"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net/http"
"os"
"path/filepath"
"regexp"
"slices"
"strings"
"sync/atomic"
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
// The actions a rule takes when it matches.
const (
// ActionLog notes the match in the request log, and does nothing else.
ActionLog = "log"
// ActionBlock refuses the request with 403.
ActionBlock = "block"
// ActionBan refuses the request and bans the client's netblock: the
// request is a clear sign of attack.
ActionBan = "ban"
)
// extension ends the name of every rule file.
const extension = ".rules"
// quietTime is how long SWWAF_RULES_DIR must go without a change before
// the rule files are read again, so that a file still being written, such
// as one saved in place, appended to or copied in with scp, is read only
// once whole.
const quietTime = 2 * time.Second
// headerTarget starts the target that is one request header,
// header:<Name>.
const headerTarget = "header:"
// escapeLength is the length of a percent escape, such as %2e.
const escapeLength = 3
var (
// ruleLine is a rule: four fields separated by spaces or tabs, of
// which the fourth, the regex, runs to the end of the line.
ruleLine = regexp.MustCompile(`^([^ \t]+)[ \t]+([^ \t]+)[ \t]+([^ \t]+)[ \t]+(.+)$`)
// idChars are the characters of a rule's id.
idChars = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
)
var (
errNotRule = errors.New(
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs")
errNotID = errors.New("is not an id of letters, digits, - and _")
errNotTarget = errors.New(
"is not path, query, uri, method, host, user_agent, referer or header:<Name>")
errNotHeaderName = errors.New(
"has a character after header: that no header name can have")
errHeaderTakenOut = errors.New(
"names a header that Go's HTTP server takes out of every request, " +
"so a rule never sees it")
errNotAction = errors.New("is not log, block or ban")
errNotRegex = errors.New("does not compile")
errUsedTwice = errors.New("is already the id of the rule at")
)
// Rule is one rule of a rule file.
type Rule struct {
// ID names the rule in the request log, the metrics and ban notes.
ID string
// Target is what the regex is matched against, such as path or
// header:Accept.
Target string
// Action is ActionLog, ActionBlock or ActionBan.
Action string
regex *regexp.Regexp
}
// Params are what Load needs.
type Params struct {
// Dir is the directory of the rule files (SWWAF_RULES_DIR).
Dir string
// Enabled is SWWAF_RULES_ENABLED: while it is false, no file is read
// and no rule loaded.
Enabled bool
// ProcessLog receives how many rules were read, and the error in a
// rule file edited while smallwebwaf runs.
ProcessLog *slog.Logger
// Alerts receive a file_error alert for that error.
Alerts *alerts.Queue
}
// Files are the rule files of a running smallwebwaf, and the rules read
// from them. They are safe for concurrent use.
type Files struct {
params Params
// rules are the rules loaded, in the order of their files' names, and
// then of their lines.
rules atomic.Pointer[[]Rule]
}
// Load reads the rules of every *.rules file in Dir, in the order of the
// files' names, unless Enabled is false. A Dir that cannot be read is an
// error, and so is a line that is not a rule, a header name with a
// character no header name can have, a rule for the Host or the
// Transfer-Encoding header, which Go's HTTP server takes out of every
// request, a regex that does not compile and an id used twice, each named
// with its file and line.
func Load(params Params) (*Files, error) {
f := &Files{params: params}
f.rules.Store(&[]Rule{})
if !params.Enabled {
return f, nil
}
rules, _, err := read(params.Dir)
if err != nil {
return nil, err
}
f.rules.Store(&rules)
f.logRead(len(rules))
return f, nil
}
// Match checks r against the rules, in order, and returns those it
// matches, up to the first whose action refuses it, block or ban, which
// is then the last one returned.
func (f *Files) Match(r *http.Request) []Rule {
var matched []Rule
for _, rule := range *f.rules.Load() {
if !rule.matches(r) {
continue
}
matched = append(matched, rule)
if rule.Action != ActionLog {
break
}
}
return matched
}
// Len returns how many rules are loaded.
func (f *Files) Len() int {
return len(*f.rules.Load())
}
// Watch watches Dir until ctx is done, and reads the rule files again
// once Dir has had no change for quietTime, after one is edited, added or
// removed, and after Watch starts watching. If they then hold an error,
// the rules stay as they were, the error is logged with its file and
// line, and the files are read again after the next change. If Dir cannot
// be watched, that is logged, and the rules stay as they were loaded.
// While Enabled is false, Watch returns at once.
func (f *Files) Watch(ctx context.Context) {
if !f.params.Enabled {
return
}
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the rule files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the rule files for edits",
"directory", f.params.Dir)
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
}
// readAfterChanges reads the rule files again once quietTime has passed
// without a change from events, until ctx is done, and logs the errors
// from errs. The wait starts at once, as if for a change, so that an edit
// saved after Load read the files, and before Dir was watched, is taken
// in too.
func (f *Files) readAfterChanges(
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
) {
quiet := time.NewTimer(quietTime)
defer quiet.Stop()
for {
select {
case <-ctx.Done():
return
case <-events:
quiet.Reset(quietTime)
case <-quiet.C:
f.readAgain()
case err := <-errs:
f.params.ProcessLog.Warn("watching the rule files failed",
"error", err.Error())
}
}
}
// readAgain reads the rule files again, in place of the rules loaded, or
// logs the error that keeps the rules as they were, and raises a
// file_error alert for it, for the file it is in.
func (f *Files) readAgain() {
rules, path, err := read(f.params.Dir)
if err != nil {
const kept = "a rule file has an error, and the rules stay as they were"
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: kept,
Detail: map[string]any{"file": path, "error": err.Error()},
})
f.params.ProcessLog.Error(kept, "error", err.Error())
return
}
f.rules.Store(&rules)
f.logRead(len(rules))
}
// logRead logs that the rule files were read, and how many rules they
// hold, which can be none.
func (f *Files) logRead(count int) {
f.params.ProcessLog.Info("read the rule files",
"directory", f.params.Dir, "rules", count)
}
// read returns the rules of every rule file in dir, in the order of the
// files' names, and then of their lines, or an error, with the path of the
// rule file it is in, or dir. A file whose name starts with a dot, such as
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's
// *.rules would not match it.
func read(dir string) ([]Rule, string, error) {
entries, err := os.ReadDir(dir)
if err != nil {
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
}
var rules []Rule
// places are where each id is, as "<file>, line <n>".
places := map[string]string{}
for _, entry := range entries {
name := entry.Name()
if entry.IsDir() || strings.HasPrefix(name, ".") || filepath.Ext(name) != extension {
continue
}
path := filepath.Join(dir, name)
rules, err = readFile(path, rules, places)
if err != nil {
return nil, path, err
}
}
return rules, "", nil
}
// readFile appends the rules of the rule file at path to rules. places
// are where each id read so far is, and gain those of the file.
func readFile(path string, rules []Rule, places map[string]string) ([]Rule, error) {
data, err := os.ReadFile(path) //nolint:gosec // a rule file, in SWWAF_RULES_DIR
if err != nil {
return nil, err
}
number := 0
for line := range strings.Lines(string(data)) {
number++
place := fmt.Sprintf("%s, line %d", path, number)
text := strings.TrimSuffix(strings.TrimSuffix(line, "\n"), "\r")
rule, isRule, err := parse(text)
if err != nil {
return nil, fmt.Errorf("%s: %w", place, err)
}
if !isRule {
continue
}
first, used := places[rule.ID]
if used {
return nil, fmt.Errorf("%s: the id %q %w %s", place, rule.ID, errUsedTwice, first)
}
places[rule.ID] = place
rules = append(rules, rule)
}
return rules, nil
}
// parse reads a line of a rule file. It returns false for a blank line
// and for a comment, a line that starts with #. Spaces and tabs at the
// end of the line are not part of its regex, so a line with only those
// after its action has no regex, and is not a rule.
func parse(line string) (Rule, bool, error) {
line = strings.Trim(line, " \t")
if line == "" || strings.HasPrefix(line, "#") {
return Rule{}, false, nil
}
fields := ruleLine.FindStringSubmatch(line)
if fields == nil {
return Rule{}, false, errNotRule
}
rule := Rule{ID: fields[1], Target: fields[2], Action: fields[3]}
headerName, isHeader := strings.CutPrefix(rule.Target, headerTarget)
switch {
case !idChars.MatchString(rule.ID):
return Rule{}, false, fmt.Errorf("the id %q %w", rule.ID, errNotID)
case !isTarget(rule.Target):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotTarget)
case isHeader && !config.IsHeaderName(headerName):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotHeaderName)
case strings.EqualFold(rule.Target, headerTarget+"Host"):
return Rule{}, false, fmt.Errorf(
"the target %q %w; the request's host is the target host",
rule.Target, errHeaderTakenOut)
case strings.EqualFold(rule.Target, headerTarget+"Transfer-Encoding"):
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errHeaderTakenOut)
case !slices.Contains([]string{ActionLog, ActionBlock, ActionBan}, rule.Action):
return Rule{}, false, fmt.Errorf("the action %q %w", rule.Action, errNotAction)
}
regex, err := regexp.Compile(fields[4])
if err != nil {
return Rule{}, false, fmt.Errorf("the regex %w: %w", errNotRegex, err)
}
rule.regex = regex
return rule, true, nil
}
// isTarget reports whether target is one a rule may have.
func isTarget(target string) bool {
switch target {
case "path", "query", "uri", "method", "host", "user_agent", "referer":
return true
}
name, isHeader := strings.CutPrefix(target, headerTarget)
return isHeader && name != ""
}
// matches reports whether the rule's regex matches its target in r. For
// uri it is matched against the path and query as received, and against
// them once percent-decoded, so that an encoded probe cannot slip past.
func (rule Rule) matches(r *http.Request) bool {
if rule.Target == "uri" {
uri := pathAndQuery(r)
return rule.regex.MatchString(uri) || rule.regex.MatchString(decodeOnce(uri))
}
return rule.regex.MatchString(value(rule.Target, r))
}
// value returns what a rule with target, other than uri, is matched
// against in r: the path and the query as the client sent them, before
// any decoding or re-encoding, split at the first ?, and a header's values
// joined by ", ", as HTTP joins those of a header sent more than once.
func value(target string, r *http.Request) string {
switch target {
case "path":
path, _, _ := strings.Cut(pathAndQuery(r), "?")
return path
case "query":
_, query, _ := strings.Cut(pathAndQuery(r), "?")
return query
case "method":
return r.Method
case "host":
return r.Host
case "user_agent":
return header(r, "User-Agent")
case "referer":
return header(r, "Referer")
default:
return header(r, strings.TrimPrefix(target, headerTarget))
}
}
// pathAndQuery returns the target of r's request line, r.RequestURI, as
// the client sent it, less any scheme and host: a target with a scheme
// gives what follows the scheme and its :, and the host when // follows.
// So http://host/path, as a client sends it to a proxy, gives /path, and
// so does http:/path, which Go reads as a target with a scheme and no
// host. r.URL is not used: when the path holds a character it escapes,
// such as \ or a non-ASCII byte, it decodes the whole path and escapes it
// again, so that \ becomes %5C and %2e a dot.
func pathAndQuery(r *http.Request) string {
if !r.URL.IsAbs() {
return r.RequestURI
}
_, afterScheme, _ := strings.Cut(r.RequestURI, ":")
hostAndRest, hasHost := strings.CutPrefix(afterScheme, "//")
if !hasHost {
return afterScheme
}
start := strings.IndexAny(hostAndRest, "/?")
if start < 0 {
return ""
}
return hostAndRest[start:]
}
// header returns the values of r's header name joined by ", ", or "" if
// r has no such header.
func header(r *http.Request, name string) string {
return strings.Join(r.Header.Values(name), ", ")
}
// decodeOnce returns s with each percent escape, such as %2e, replaced by
// the byte it stands for. A % that is not followed by two hex digits is
// left as it is, so that a malformed escape cannot keep the rest of s
// from being decoded.
func decodeOnce(s string) string {
var decoded strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '%' && i+escapeLength <= len(s) {
b, err := hex.DecodeString(s[i+1 : i+escapeLength])
if err == nil {
decoded.Write(b)
i += escapeLength - 1
continue
}
}
decoded.WriteByte(s[i])
}
return decoded.String()
}
+687
View File
@@ -0,0 +1,687 @@
package rules_test
import (
"context"
"encoding/json"
"log/slog"
"maps"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"slices"
"strconv"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/rules"
)
const (
// What the process log says once Watch watches the directory, after
// each reading of the rule files, and for one that has an error.
watching = "watching the rule files for edits"
read = "read the rule files"
hasError = "a rule file has an error, and the rules stay as they were"
// maxLogLines is how many lines of the process log wait for a test to
// read them.
maxLogLines = 64
// browser is the user agent of an ordinary visitor.
browser = "Mozilla/5.0 (X11; Linux x86_64; rv:140.0) Gecko/20100101 Firefox/140.0"
// testFile is the rule file of a test that needs only one, and
// firstFile the first of a test's rule files.
testFile = "test.rules"
firstFile = "00-a.rules"
// userAgent is the header that carries the user agent.
userAgent = "User-Agent"
)
func TestEachTargetMatchesWhatItNames(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
rule string // its target, action and regex
uri string // the request's path and query
header http.Header
want bool
}{
{"path as received", `path log ^/%2eenv$`, "/%2eenv", nil, true},
{"path not decoded", `path log ^/\.env$`, "/%2eenv", nil, false},
{"path without the query", `path log ^/a$`, "/a?b=c", nil, true},
{"query as received", `query log ^b=%2e$`, "/a?b=%2e", nil, true},
{"uri as received", `uri log ^/a\?b=%2e$`, "/a?b=%2e", nil, true},
{"uri decoded", `uri log (\.\./){2}`, "/a?f=%2e%2e%2f%2e%2e%2f", nil, true},
{
"uri decoded past malformed escapes", `uri log (\.\./){2}&h=%$`,
"/a?g=%zz&f=%2e%2e%2f%2e%2e%2f&h=%", nil, true,
},
{"uri decoded only once", `uri log ^/a\.b$`, "/a%252eb", nil, false},
{"method", `method log ^PUT$`, "/", nil, true},
{"host", `host log ^app\.example$`, "/", nil, true},
{
"user_agent", `user_agent log ^sqlmap/`, "/",
http.Header{userAgent: {"sqlmap/1.8"}}, true,
},
{
"user_agent sent twice", `user_agent log ^curl/8, sqlmap/`, "/",
http.Header{userAgent: {"curl/8", "sqlmap/1.8"}}, true,
},
{"user_agent missing", `user_agent log ^$`, "/", nil, true},
{
"referer", `referer log ^https://spam\.example/`, "/",
http.Header{"Referer": {"https://spam.example/buy"}}, true,
},
{
"a header sent twice", `header:x-api-version log ^2, 3$`, "/",
http.Header{"X-Api-Version": {"2", "3"}}, true,
},
{"a header missing", `header:X-Api-Version log ^$`, "/", nil, true},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: "a-rule " + tc.rule + "\n"})
// Every request is a PUT, which the method rule looks for.
r := httptest.NewRequestWithContext(t.Context(), http.MethodPut,
"http://app.example"+tc.uri, nil)
maps.Copy(r.Header, tc.header)
got := len(files.Match(r)) == 1
if got != tc.want {
t.Errorf("%s matches %s: %t, want %t", tc.rule, tc.uri, got, tc.want)
}
})
}
}
func TestPathMatchedAsTheClientSentIt(t *testing.T) {
t.Parallel()
// Each path holds a character Go's URL type would escape again, \ or
// a non-ASCII byte, and each rule is written for the path as sent.
for _, tc := range []struct {
rule string // its target, action and regex
sent string // the path and query the client sent
}{
{`path log ^/\.\.\\\.\.\\windows\\win\.ini$`, `/..\..\windows\win.ini`},
{`path log ^/%2e%2e\\%2e%2e\\windows\\win\.ini$`, `/%2e%2e\%2e%2e\windows\win.ini`},
{`path log ^/café$`, "/café?x=1"},
{`uri log ^/%2e%2e\\%2e%2e\\boot\.ini\?x=1$`, `/%2e%2e\%2e%2e\boot.ini?x=1`},
} {
files := load(t, ruleFiles{testFile: "as-sent " + tc.rule + "\n"})
// The target in origin form, as traefik sends it, in absolute form,
// as a client sends it to a proxy, and with a scheme but no host,
// which Go reads as absolute form with no host, sending the app
// the path.
for _, target := range []string{
tc.sent, "http://app.example" + tc.sent, "http:" + tc.sent, "foo:" + tc.sent,
} {
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet, target, nil)
wantMatched(t, files, r, "as-sent")
}
}
}
func TestMatchingStopsAtTheFirstRuleThatRefuses(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: `
every-path path log ^/
no-path path log ^$
first-refusal path block ^/probe
later-ban path ban ^/probe
after path log ^/
`})
// Every log rule that matches is noted, and the block rule ends the
// matching.
wantMatched(t, files, get(t, "/probe"), "every-path", "first-refusal")
wantMatched(t, files, get(t, "/page"), "every-path", "after")
// A ban rule ends it too.
files = load(t, ruleFiles{testFile: "ban path ban ^/\nlater path block ^/\n"})
wantMatched(t, files, get(t, "/"), "ban")
}
func TestSpacesAndTabsEndingALineAreNotPartOfItsRegex(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{testFile: "env-file path block ^/\\.env$ \t \n"})
wantMatched(t, files, get(t, "/.env"), "env-file")
}
func TestFilesReadInNameOrderThenLineOrder(t *testing.T) {
t.Parallel()
files := load(t, ruleFiles{
"50-b.rules": "b1 path log ^/\n\n# a comment\n # an indented one\nb2 path log ^/\n",
firstFile: "a1 path log ^/\r\n",
// None is a rule file.
"notes.txt": "notes, not rules\n",
"10-c.rules.bak": "an old copy\n",
"20-d.rules/keep": "a file in a directory\n",
})
wantMatched(t, files, get(t, "/"), "a1", "b1", "b2")
if files.Len() != 3 {
t.Errorf("%d rules loaded, want 3", files.Len())
}
}
func TestFileWhoseNameStartsWithADotIsNotARuleFile(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "probe path block ^/probe\n"})
// The lock file Emacs makes beside a file while it is edited: a link to
// nothing, which cannot be read.
err := os.Symlink("user@host.1234:1700000000", filepath.Join(dir, ".#"+firstFile))
if err != nil {
t.Fatalf("symlink: %v", err)
}
params, _ := newParams(dir)
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
wantMatched(t, files, get(t, "/probe"), "probe")
}
func TestFaultStopsTheStartNamingTheFileAndLine(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
content string
line int
want string
}{
{
"too few fields", "env-file path ban\n", 1,
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs",
},
{
// Else its regex would be a space, found in nearly every user agent.
"a regex of only spaces and tabs", "scanner user_agent ban\t \n", 1,
"is not a rule: an id, a target, an action and a regex, " +
"separated by spaces or tabs",
},
{
"an id of other characters", "# ids\n\nenv.file path ban ^/\n", 3,
`the id "env.file" is not an id of letters, digits, - and _`,
},
{
"an unknown target", "env-file paths ban ^/\n", 1,
`the target "paths" is not path, query, uri, method, host, ` +
"user_agent, referer or header:<Name>",
},
{
"a header without a name", "env-file header: ban ^/\n", 1,
`the target "header:" is not path, query, uri, method, host, ` +
"user_agent, referer or header:<Name>",
},
{
"a header name written with its colon", "sqlmap header:User-Agent: ban sqlmap\n", 1,
`the target "header:User-Agent:" has a character after header: ` +
"that no header name can have",
},
{
"a header name with a semicolon", "accept header:Accept;q log ^$\n", 1,
`the target "header:Accept;q" has a character after header: ` +
"that no header name can have",
},
{
"a header name with brackets", "x-header header:X(y) log ^$\n", 1,
`the target "header:X(y)" has a character after header: ` +
"that no header name can have",
},
{
"the Host header", "host-header header:host block ^$\n", 1,
`the target "header:host" names a header that Go's HTTP server ` +
"takes out of every request, so a rule never sees it; " +
"the request's host is the target host",
},
{
"the Transfer-Encoding header",
"# bodies sent in chunks\nchunked header:Transfer-Encoding block ^chunked$\n", 2,
`the target "header:Transfer-Encoding" names a header that Go's ` +
"HTTP server takes out of every request, so a rule never sees it",
},
{
"an unknown action", "env-file path deny ^/\n", 1,
`the action "deny" is not log, block or ban`,
},
{
"a regex that does not compile", "env-file path ban ^/(\n", 1,
"the regex does not compile: error parsing regexp: " +
"missing closing ): `^/(`",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{"00-default.rules": tc.content})
path := filepath.Join(dir, "00-default.rules")
wantRefused(t, dir, path+", line "+strconv.Itoa(tc.line)+": "+tc.want)
})
}
}
func TestIDUsedTwiceStopsTheStartNamingBothPlaces(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{
"00-a.rules": "probe path log ^/a\n",
"50-b.rules": "other path log ^/b\nprobe path ban ^/c\n",
})
wantRefused(t, dir, filepath.Join(dir, "50-b.rules")+`, line 2: the id "probe" `+
"is already the id of the rule at "+filepath.Join(dir, "00-a.rules")+", line 1")
}
func TestDirectoryThatDoesNotExistStopsTheStart(t *testing.T) {
t.Parallel()
dir := filepath.Join(t.TempDir(), "rules.d")
wantRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
": no such file or directory")
}
func TestEmptyDirectoryLoadsNoRulesAndSaysSo(t *testing.T) {
t.Parallel()
params, lines := newParams(writeFiles(t, ruleFiles{"00-default.rules": "# none\n"}))
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
line := lines.waitFor(t, read)
if files.Len() != 0 || line["rules"] != 0.0 {
t.Errorf("%d rules loaded, and the log says %v, want none", files.Len(), line)
}
}
func TestRuleFilesOffReadNothing(t *testing.T) {
t.Parallel()
// SWWAF_RULES_DIR does not exist, which would stop the start.
params, _ := newParams(filepath.Join(t.TempDir(), "rules.d"))
params.Enabled = false
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
if files.Len() != 0 || files.Match(get(t, "/")) != nil {
t.Errorf("%d rules loaded with the rule files off", files.Len())
}
// It would watch until the test ends.
files.Watch(t.Context())
}
func TestEditsTakenInWhileRunning(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines, _ := watch(t, dir)
// matches reports whether path matches a rule.
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
// A file added.
save(t, dir, "50-b.rules", "second path block ^/second\n")
lines.waitUntil(t, func() bool { return matches("/second") })
wantMatched(t, files, get(t, "/first"), "first")
// A file edited.
save(t, dir, firstFile, "first path block ^/edited\n")
lines.waitUntil(t, func() bool { return !matches("/first") })
wantMatched(t, files, get(t, "/edited"), "first")
// A file removed.
err := os.Remove(filepath.Join(dir, "50-b.rules"))
if err != nil {
t.Fatalf("remove: %v", err)
}
lines.waitUntil(t, func() bool { return !matches("/second") })
wantMatched(t, files, get(t, "/edited"), "first")
}
func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines, queue := watch(t, dir)
// The edit's second line has an unknown action, so the rules stay as
// they were, the first line's earlier version included.
save(t, dir, firstFile, "first path block ^/edited\nsecond path bann ^/second\n")
line := lines.waitFor(t, hasError)
want := filepath.Join(dir, firstFile) +
`, line 2: the action "bann" is not log, block or ban`
if line["error"] != want || line["level"] != "ERROR" {
t.Errorf("logged %v, want an error %q", line, want)
}
// The error is raised as a file_error alert too, for the file.
wantFileError := func() {
t.Helper()
waiting := queue.Snapshot().Waiting
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
waiting[0].Reason != hasError || waiting[0].Detail["error"] != want ||
waiting[0].Detail["file"] != filepath.Join(dir, firstFile) {
t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want)
}
}
wantFileError()
wantMatched(t, files, get(t, "/first"), "first")
wantMatched(t, files, get(t, "/second"))
// Once mended, the file is read again, and raises no alert.
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
wantMatched(t, files, get(t, "/edited"), "first")
wantFileError()
}
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
t.Parallel()
params, _ := newParams(filepath.Join("..", "..", "share", "rules.d"))
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load the default file: %v", err)
}
// Probes sent by a browser, by the rule that refuses them.
for rule, targets := range map[string][]string{
"env-file": {"/.env", "/.env.production", "/.ENV"},
"vcs-dir": {"/.git/config", "/.git", "/.svn/entries"},
"secrets-dir": {"/.aws/credentials", "/.ssh/id_rsa"},
"secret-file": {"/.htpasswd", "/.DS_Store", "/.git-credentials"},
"editor-dir": {"/.vscode/sftp.json"},
"backup-file": {
"/wp-config.php.bak", "/index.php~", "/dump.sql", "/backup.sql.gz",
},
"log-file": {"/debug.log"},
"compose-file": {"/docker-compose.yml", "/compose.yaml"},
"php-shell": {"/shell.php"},
"path-traversal": {
"/static/../../etc/passwd", "/f?f=%2e%2e%2f%2e%2e%2fetc%2fpasswd",
},
} {
for _, target := range targets {
wantRefusedBy(t, files, target, browser, rule)
}
}
// Scanners, by their user agents.
for _, scanner := range []string{
"sqlmap/1.8.4#stable (https://sqlmap.org)",
"Mozilla/5.0 (compatible; Nuclei - Open-source project)",
} {
wantRefusedBy(t, files, "/", scanner, "scanner-agent")
}
// Ordinary requests to a code forge for files of those names deeper
// in its paths, and for other files at its root.
for _, target := range []string{
"/owner/repo/src/branch/main/.env.example",
"/owner/repo/src/branch/main/.env",
"/owner/repo/src/branch/main/.github/workflows/ci.yml",
"/owner/repo/src/branch/main/.vscode/settings.json",
"/owner/repo/src/branch/main/.htaccess",
"/owner/repo/src/branch/main/docker-compose.yml",
"/owner/repo/src/branch/main/db/schema.sql",
"/owner/repo/raw/branch/main/debug.log",
"/owner/repo.git/info/refs?service=git-upload-pack",
"/owner/repo/src/branch/main/docs/../README.md",
"/user/login?redirect_to=%2fowner%2frepo",
"/index.php",
"/.well-known/security.txt",
} {
r := get(t, target)
r.Header.Set(userAgent, browser)
matched := files.Match(r)
if len(matched) != 0 {
t.Errorf("%s matched %v, want no rule", target, ids(matched))
}
}
// A request without a user agent is only noted.
wantMatched(t, files, get(t, "/"), "empty-agent")
}
// ruleFiles are files to write into a directory of rule files, by name.
type ruleFiles map[string]string
// writeFiles writes files into a new directory, and returns it.
func writeFiles(t *testing.T, files ruleFiles) string {
t.Helper()
dir := t.TempDir()
for name, content := range files {
path := filepath.Join(dir, name)
err := os.MkdirAll(filepath.Dir(path), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
err = os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
}
return dir
}
// save writes content to the rule file name in dir as an editor that
// saves by renaming does, so that the file is never seen half written.
func save(t *testing.T, dir, name, content string) {
t.Helper()
path := filepath.Join(dir, name)
err := os.WriteFile(path+".tmp", []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", name, err)
}
err = os.Rename(path+".tmp", path)
if err != nil {
t.Fatalf("rename: %v", err)
}
}
// newParams returns Params for the rule files in dir, switched on, with
// the process log in the processLog returned, and the alerts waiting in a
// queue for a webhook that is never sent them.
func newParams(dir string) (rules.Params, processLog) {
lines := make(processLog, maxLogLines)
return rules.Params{
Dir: dir,
Enabled: true,
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
Alerts: alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: time.Now,
}),
}, lines
}
// load writes files into a new directory and loads the rules in it.
func load(t *testing.T, files ruleFiles) *rules.Files {
t.Helper()
params, _ := newParams(writeFiles(t, files))
params.ProcessLog = slog.New(slog.DiscardHandler)
loaded, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
return loaded
}
// watch loads the rules in dir, runs their Watch until the test ends, and
// waits until it watches the directory. It returns the alerts' queue as
// well.
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
t.Helper()
params, lines := newParams(dir)
files, err := rules.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Watch(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
lines.waitFor(t, watching)
return files, lines, params.Alerts
}
// wantRefused checks that loading the rule files in dir fails with the
// error want.
func wantRefused(t *testing.T, dir, want string) {
t.Helper()
params, _ := newParams(dir)
_, err := rules.Load(params)
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
// get returns a GET request for target, a path and an optional query, as
// smallwebwaf's server reads it, without a user agent.
func get(t *testing.T, target string) *http.Request {
t.Helper()
return httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"http://app.example"+target, nil)
}
// wantRefusedBy checks that a GET request for target with the user agent
// sent matches rule alone, and that rule refuses it.
func wantRefusedBy(t *testing.T, files *rules.Files, target, sent, rule string) {
t.Helper()
r := get(t, target)
r.Header.Set(userAgent, sent)
matched := files.Match(r)
if len(matched) != 1 || matched[0].ID != rule || matched[0].Action == rules.ActionLog {
t.Errorf("%s from %q matched %v, want %s alone, refusing it", target,
sent, ids(matched), rule)
}
}
// wantMatched checks the ids of the rules r matches, in order.
func wantMatched(t *testing.T, files *rules.Files, r *http.Request, want ...string) {
t.Helper()
got := ids(files.Match(r))
if !slices.Equal(got, want) {
t.Errorf("%s matched %v, want %v", r.URL, got, want)
}
}
// ids returns the ids of matched.
func ids(matched []rules.Rule) []string {
got := make([]string, 0, len(matched))
for _, rule := range matched {
got = append(got, rule.ID)
}
return got
}
// processLog receives the lines of a process log, each a JSON object, for
// a test to wait for.
type processLog chan string
// Write receives a line of the process log.
func (l processLog) Write(line []byte) (int, error) {
l <- string(line)
return len(line), nil
}
// waitFor returns the next line of the process log whose message is msg,
// passing over the lines before it. It waits as long as that takes, so
// that a slow test process cannot fail the test.
func (l processLog) waitFor(t *testing.T, msg string) map[string]any {
t.Helper()
for line := range l {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil {
t.Fatalf("process log line %q is not JSON: %v", line, err)
}
if fields["msg"] == msg {
return fields
}
}
return nil
}
// waitUntil waits for the rule files to be read until done reports true,
// as it does once they have been read after the test's last change. They
// can be read before then too, as they are once Watch starts watching.
func (l processLog) waitUntil(t *testing.T, done func() bool) {
t.Helper()
for !done() {
l.waitFor(t, read)
}
}
+165
View File
@@ -0,0 +1,165 @@
package rules
import (
"context"
"log/slog"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"slices"
"testing"
"testing/synctest"
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// The tests below run readAfterChanges in a synctest bubble, where time is
// a clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once readAfterChanges waits again, so that every
// reading due by then is done. The test sends the changes itself, as the
// watch of a directory cannot run in a bubble.
func TestFileWrittenInTwoPartsTakenInOnlyWhole(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "50-app.rules")
writeFile(t, path, "first path block ^/first\n")
files := load(t, dir)
changes := run(t, files)
file, err := os.Create(path) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("create: %v", err)
}
defer func() {
_ = file.Close()
}()
// The first part ends in the middle of a ban rule's regex, which,
// read then, would ban every request.
write(t, file, "first path block ^/first\nprobe path ban ^/")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/anything")
// The second part starts the wait again.
write(t, file, `\.env$`+"\n")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/.env")
time.Sleep(time.Nanosecond)
synctest.Wait()
wantMatched(t, files, "/.env", "probe")
wantMatched(t, files, "/anything")
})
}
func TestEditSavedBeforeTheWatchStartsTakenIn(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "50-app.rules")
writeFile(t, path, "first path block ^/first\n")
files := load(t, dir)
// Saved after Load read the files, and before the directory was
// watched, so that no change is seen for it.
writeFile(t, path, "first path block ^/edited\n")
run(t, files)
time.Sleep(quietTime)
synctest.Wait()
wantMatched(t, files, "/edited", "first")
})
}
// load loads the rules in dir.
func load(t *testing.T, dir string) *Files {
t.Helper()
files, err := Load(Params{
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
Alerts: alerts.New(alerts.Params{}),
})
if err != nil {
t.Fatalf("load: %v", err)
}
return files
}
// run runs files' readAfterChanges until the test ends, and returns the
// channel that sends it changes.
func run(t *testing.T, files *Files) chan<- fsnotify.Event {
t.Helper()
changes := make(chan fsnotify.Event)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.readAfterChanges(ctx, changes, nil)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
return changes
}
// writeFile writes content to the file at path.
func writeFile(t *testing.T, path, content string) {
t.Helper()
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
}
// write writes text to the end of file.
func write(t *testing.T, file *os.File, text string) {
t.Helper()
_, err := file.WriteString(text)
if err != nil {
t.Fatalf("write: %v", err)
}
}
// wantMatched checks the ids of the rules that a GET request for path
// matches, in order.
func wantMatched(t *testing.T, files *Files, path string, want ...string) {
t.Helper()
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"http://app.example"+path, nil)
matched := files.Match(r)
got := make([]string, 0, len(matched))
for _, rule := range matched {
got = append(got, rule.ID)
}
if !slices.Equal(got, want) {
t.Errorf("%s matched %v, want %v", path, got, want)
}
}
+5 -3
View File
@@ -23,6 +23,8 @@ var errHealthEndpoint = errors.New("smallwebwaf's health endpoint answered")
// smallwebwaf answers its health endpoint on 127.0.0.1, at the port in
// SWWAF_LISTEN_ADDR, and the app accepts connections at the address in
// SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1.
// It reads no other setting, nor a file that another names, so neither
// can fail it.
// args are the arguments after `healthcheck`; it takes none, and given
// one it names it on stderr and returns 1 without checking anything.
func HealthCheck(
@@ -50,13 +52,13 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout)
defer cancel()
cfg, err := config.FromEnvironment(lookupEnv)
listenAddr, upstreamURL, err := config.ListenAddrAndUpstreamURL(lookupEnv)
if err != nil {
return fmt.Errorf("invalid setting: %w", err)
}
// The settings have checked that the address has a port.
_, port, _ := net.SplitHostPort(cfg.ListenAddr)
_, port, _ := net.SplitHostPort(listenAddr)
health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath
req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody)
@@ -75,7 +77,7 @@ func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) err
return fmt.Errorf("%w %s", errHealthEndpoint, res.Status)
}
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(cfg.UpstreamURL))
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(upstreamURL))
if err != nil {
return fmt.Errorf("connect to the app: %w", err)
}
+33
View File
@@ -6,6 +6,8 @@ import (
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
@@ -28,6 +30,7 @@ func TestHealthCheck(t *testing.T) {
listenAddr: localhost + ":0",
upstreamURL: app.URL,
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
}
go func() {
@@ -42,6 +45,21 @@ func TestHealthCheck(t *testing.T) {
wantHealthCheck(t, env, 0, "")
// The health check reads those two settings alone, here given as
// files: a removed or invalid token file, or an invalid value of
// another setting, does not fail it.
for _, other := range []struct{ name, value string }{
{"SWWAF_METRICS_TOKEN_FILE", filepath.Join(t.TempDir(), "removed")},
{"SWWAF_METRICS_TOKEN_FILE", writeFile(t, "too short\n")},
{"SWWAF_MODE", "neither"},
} {
wantHealthCheck(t, map[string]string{
listenAddr + "_FILE": writeFile(t, ":"+port+"\n"),
upstreamURL + "_FILE": writeFile(t, app.URL+"\n"),
other.name: other.value,
}, 0, "")
}
app.Close()
wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ")
@@ -97,3 +115,18 @@ func wantHealthCheck(t *testing.T, env map[string]string, status int, message st
got, wrote, status, message)
}
}
// writeFile writes contents to a file in a directory of its own, removed
// when the test ends, and returns the file's path.
func writeFile(t *testing.T, contents string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "setting")
err := os.WriteFile(path, []byte(contents), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
return path
}
+157 -37
View File
@@ -1,6 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings
// and the state files, serves requests until it is told to stop, and then
// stops in an orderly way, writing the state files.
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// the rule files and the state files, serves requests until it is told to
// stop, and then stops in an orderly way, writing the state files.
package smallwebwaf
import (
@@ -15,10 +15,13 @@ import (
"syscall"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"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/rules"
"sneak.berlin/go/smallwebwaf/internal/state"
)
@@ -27,6 +30,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.
@@ -56,9 +64,9 @@ func Main(version string) int {
})
}
// Run reads the settings and the state files, then serves requests until
// ctx is done. It returns the process's exit status, 1 when smallwebwaf
// cannot start.
// Run reads the settings, the rule files and the state files, then serves
// requests until ctx is done. It returns the process's exit status, 1
// when smallwebwaf cannot start.
func Run(ctx context.Context, params Params) int {
processLog := requestlog.NewProcessLogger(params.Stdout)
@@ -69,28 +77,56 @@ func Run(ctx context.Context, params Params) int {
return 1
}
// The state files give times in UTC.
// While SWWAF_LOG_REMOTE_URL is set, every line on stdout from here on
// is sent there too.
stdout := params.Stdout
var remote *remotelog.Sender
if cfg.LogRemoteURL != nil {
remote = newRemoteLogSender(cfg)
stdout = io.MultiWriter(params.Stdout, remote)
processLog = requestlog.NewProcessLogger(stdout)
stopSending := startSending(ctx, remote, processLog)
defer stopSending()
}
// The state files and the alerts give times in UTC.
now := func() time.Time { return time.Now().UTC() }
alertQueue := newAlertQueue(cfg, now, processLog)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir,
Enabled: cfg.RulesEnabled,
ProcessLog: processLog,
Alerts: alertQueue,
})
if err != nil {
processLog.Error("cannot use the rule files", "error", err.Error())
return 1
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
RequestLog: stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
})
if remote != nil {
server.Metrics.AddRemoteLog(remote)
}
files, err := state.Load(state.Params{
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
if cfg.AlertWebhookURL != nil {
server.Metrics.AddAlerts(alertQueue)
}
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
if err != nil {
processLog.Error("cannot use the state files", "error", err.Error())
@@ -110,16 +146,92 @@ func Run(ctx context.Context, params Params) int {
"address", listener.Addr().String(),
"settings", cfg)
return serve(ctx, server.Server, listener, files, processLog)
return serve(ctx, server.Server, listener, files, ruleFiles, alertQueue, processLog)
}
// loadStateFiles reads the state files into the parts of server and into
// alertQueue, as state.Load does.
func loadStateFiles(
cfg *config.Config, server *proxy.Server, alertQueue *alerts.Queue,
now func() time.Time, processLog *slog.Logger,
) (*state.Files, error) {
return state.Load(state.Params{
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Alerts: alertQueue,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
}
// newAlertQueue returns the queue of the alerts to SWWAF_ALERT_WEBHOOK_URL,
// with the settings for it.
func newAlertQueue(
cfg *config.Config, now func() time.Time, processLog *slog.Logger,
) *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
WebhookHeaders: cfg.AlertWebhookHeaders,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
}
// newRemoteLogSender returns a sender of the log lines to
// SWWAF_LOG_REMOTE_URL, with the settings for it.
func newRemoteLogSender(cfg *config.Config) *remotelog.Sender {
return remotelog.New(remotelog.Params{
URL: cfg.LogRemoteURL,
RootCAs: cfg.LogRemoteTLSCAs,
Buffer: cfg.LogRemoteBuffer,
Facility: cfg.LogRemoteFacility,
AppName: cfg.LogRemoteAppName,
})
}
// startSending runs remote until the function it returns is called, which
// then waits at most remoteLogStopTimeout for the lines still waiting to
// be sent. Sending goes on after ctx is done, so that the lines written
// while smallwebwaf stops are sent too.
func startSending(
ctx context.Context, remote *remotelog.Sender, processLog *slog.Logger,
) func() {
sending, stop := context.WithCancel(context.WithoutCancel(ctx))
sent := make(chan struct{})
go func() {
remote.Run(sending, processLog)
close(sent)
}()
return func() {
stop()
select {
case <-sent:
case <-time.After(remoteLogStopTimeout):
}
}
}
// serve serves requests on listener, writes the state files as they are
// due, 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.
// due, takes in an admin's edits of them, reads the rule files again as
// they change, and sends the alerts, until ctx is done. Then it gives the
// requests in progress shutdownTimeout to finish, and writes every state
// file, alerts.json with the alerts still waiting.
func serve(
ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, processLog *slog.Logger,
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
processLog *slog.Logger,
) int {
served := make(chan error, 1)
@@ -130,18 +242,10 @@ func serve(
writing, stopWriting := context.WithCancel(ctx)
defer stopWriting()
written := make(chan struct{})
watched := make(chan struct{})
go func() {
files.Run(writing)
close(written)
}()
go func() {
files.Watch(writing)
close(watched)
}()
written := inBackground(func() { files.Run(writing) })
watched := inBackground(func() { files.Watch(writing) })
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
alertsSent := inBackground(func() { alertQueue.Run(writing) })
select {
case err := <-served:
@@ -173,7 +277,8 @@ func serve(
}
// Run and Watch have ended, so nothing else reads or writes the
// files. Every request has ended too, but for two kinds
// files, and no alert is being sent, so that alerts.json keeps every
// alert not yet sent. 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
@@ -181,6 +286,8 @@ func serve(
// missing from clients.json.
<-written
<-watched
<-rulesWatched
<-alertsSent
err = files.WriteAll()
if err != nil {
@@ -193,3 +300,16 @@ func serve(
return 0
}
// inBackground runs task on a goroutine of its own, and returns a channel
// that is closed once task has returned.
func inBackground(task func()) <-chan struct{} {
done := make(chan struct{})
go func() {
task()
close(done)
}()
return done
}
@@ -0,0 +1,82 @@
package smallwebwaf
import (
"log/slog"
"net/url"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/remotelog"
)
// The stop's tests run in a synctest bubble, where the time package runs
// on a clock of the test's own, so that how long the stop takes can be
// told exactly. The sender is held up by its process log, not by the
// network: a goroutine of the bubble that waits on the network keeps that
// clock from moving on.
func TestStopWaitsForTheSenderToFinish(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
took := stopHeldSender(t, time.Second)
if took != time.Second {
t.Errorf("the stop took %s, want the second the sender took", took)
}
})
}
func TestStopWaitsForTheSenderAtMostTwoSeconds(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
took := stopHeldSender(t, time.Minute)
if took != 2*time.Second {
t.Errorf("the stop took %s, want 2s", took)
}
})
}
// heldLog holds each line written to it until it is closed.
type heldLog chan struct{}
// Write waits until the log is closed.
func (l heldLog) Write(p []byte) (int, error) {
<-l
return len(p), nil
}
// stopHeldSender starts sending to an endpoint the sender cannot connect
// to, holds the sender as it logs that failure until release has passed,
// stops the sending, and returns how long the stop took. It returns once
// the sender has ended, as a bubble must.
func stopHeldSender(t *testing.T, release time.Duration) time.Duration {
t.Helper()
log := make(heldLog)
sender := remotelog.New(remotelog.Params{
// No port is 65536, so each attempt to connect fails at once,
// before it reaches the network.
URL: &url.URL{Scheme: remotelog.SchemeTCP, Host: "127.0.0.1:65536"},
Buffer: 1,
})
stopSending := startSending(t.Context(), sender,
slog.New(slog.NewJSONHandler(log, nil)))
synctest.Wait()
time.AfterFunc(release, func() { close(log) })
stopped := time.Now()
stopSending()
took := time.Since(stopped)
time.Sleep(release)
synctest.Wait()
return took
}
+553 -13
View File
@@ -10,8 +10,11 @@ import (
"net/http/httptest"
"os"
"path/filepath"
"slices"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
@@ -34,6 +37,10 @@ const (
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rulesDir = "SWWAF_RULES_DIR"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set.
adminSecret = "fedcba9876543210fedcba9876543210"
// greeting is what the tests' app answers.
greeting = "hello from the app"
)
@@ -124,25 +131,31 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
}
}
func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) {
func TestShortTokenStopsTheStartUnshown(t *testing.T) {
t.Parallel()
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
out := &output{}
for _, name := range []string{adminToken, "SWWAF_METRICS_TOKEN"} {
t.Run(name, func(t *testing.T) {
t.Parallel()
status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
}
out := &output{}
line := out.line(t, "msg", "invalid setting")
if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" {
t.Errorf("start refused with %v", line)
}
status := run(t.Context(), map[string]string{name: token}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
}
if strings.Contains(out.text(), token) {
t.Errorf("the output shows the token:\n%s", out.text())
line := out.line(t, "msg", "invalid setting")
if line["error"] != name+": is shorter than 32 characters" {
t.Errorf("start refused with %v", line)
}
if strings.Contains(out.text(), token) {
t.Errorf("the output shows the token:\n%s", out.text())
}
})
}
}
@@ -163,6 +176,7 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
status := run(t.Context(), map[string]string{
listenAddr: taken.Addr().String(),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
@@ -186,9 +200,15 @@ func TestServesUntilToldToStop(t *testing.T) {
listenAddr: localhost + ":0",
upstreamURL: appURL,
stateDir: dir,
rulesDir: filepath.Join("..", "..", "share", "rules.d"),
}, out)
}()
// The default rule file is read.
if rules := out.line(t, "msg", "read the rule files")["rules"]; rules != 12.0 {
t.Errorf("read %v rules from the default rule file, want 12", rules)
}
starting := out.line(t, "msg", "starting")
wantStartingLine(t, starting, appURL, dir)
@@ -217,6 +237,7 @@ func TestStateKeptAcrossRestarts(t *testing.T) {
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
rateLimitPerDay: "2",
// Neither comes due in the test: the files are written as
// smallwebwaf stops.
@@ -253,6 +274,7 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
rateLimitPerDay: "1",
scope: "24",
@@ -300,6 +322,7 @@ func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) {
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
// No write comes due in the test, so only the watch on the
// directory can take the edits in.
@@ -316,6 +339,306 @@ func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) {
})
}
func TestBanAddedAndLiftedThroughTheEndpointsKeptInBansJSON(t *testing.T) {
t.Parallel()
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
adminToken: adminSecret,
// Neither comes due in the test: the files are written as
// smallwebwaf stops.
stateWriteDelay: "1h",
stateCounterInterval: "1h",
}
runUntilStopped(t, env, func(url string) {
askAsAdmin(t, http.MethodPost, url+"_smallwebwaf/bans",
`{"netblock": "203.0.113.0/24", "duration": "permanent", `+
`"reason": "probes for logins"}`)
wantStatus(t, url, "203.0.113.9", http.StatusForbidden)
})
ban := onlyBan(t, dir)
if ban["netblock"] != "203.0.113.0/24" || ban["cause"] != "admin" ||
ban["reason"] != "probes for logins" || ban["expires"] != nil ||
ban["lifted"] != nil {
t.Errorf("bans.json holds %v, want the admin's permanent ban", ban)
}
// After a restart the ban still refuses; once lifted, it refuses no
// more, and bans.json keeps it, marked lifted.
runUntilStopped(t, env, func(url string) {
wantStatus(t, url, "203.0.113.9", http.StatusForbidden)
askAsAdmin(t, http.MethodDelete, url+"_smallwebwaf/bans/203.0.113.9", "")
wantStatus(t, url, "203.0.113.9", http.StatusOK)
})
ban = onlyBan(t, dir)
if ban["netblock"] != "203.0.113.0/24" || ban["lifted"] == nil {
t.Errorf("bans.json holds %v, want the admin's ban, lifted", ban)
}
}
func TestRuleFileAddedWhileRunningTakesEffect(t *testing.T) {
t.Parallel()
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: dir,
// The requests sent until the rule takes effect must not break a
// rate limit, whose ban would refuse them too.
"SWWAF_RATE_LIMIT_PER_MINUTE": "off",
}
out := runUntilStopped(t, env, func(url string) {
wantGreeting(t, url)
// Written once: each change would start the rule files' wait
// again. A file written before smallwebwaf watches the directory is
// read once it does.
err := os.WriteFile(filepath.Join(dir, "50-app.rules"),
[]byte("everything path block ^/\n"), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
// As long as that takes, so that a slow test process cannot fail
// the test.
for statusFrom(t, url, "203.0.113.9") != http.StatusForbidden {
time.Sleep(pollInterval)
}
})
out.line(t, "action", "rule_blocked")
}
func TestRuleFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "00-default.rules")
err := os.WriteFile(path, []byte("# probes\nenv-file path bann ^/\\.env$\n"), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
wantRulesRefused(t, dir, path+`, line 2: the action "bann" is not log, block or ban`)
}
func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
t.Parallel()
dir := filepath.Join(t.TempDir(), "rules.d")
wantRulesRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
": no such file or directory")
}
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
t.Parallel()
endpoint, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
_ = endpoint.Close()
}()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
"SWWAF_LOG_REMOTE_URL": "syslog+tcp://" + endpoint.Addr().String(),
}
out := runUntilStopped(t, env, func(url string) {
wantGreeting(t, url)
})
out.line(t, "type", "request")
// smallwebwaf connected as it started, and closes the connection once
// it has sent the lines written as it stopped.
conn, err := endpoint.Accept()
if err != nil {
t.Fatalf("accept: %v", err)
}
received, err := io.ReadAll(conn)
_ = conn.Close()
if err != nil {
t.Fatalf("read: %v", err)
}
// Lines written at once by several goroutines may reach stdout and
// the endpoint in different orders.
sent := messages(t, string(received))
written := slices.Collect(strings.Lines(out.text()))
slices.Sort(sent)
slices.Sort(written)
if !slices.Equal(sent, written) {
t.Errorf("sent\n%v\nwrote\n%v", sent, written)
}
}
func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
t.Parallel()
const token = "0123456789abcdef0123456789abcdef"
// The endpoint takes connections and never answers, so the TLS
// handshake of each waits on it, and no line is ever sent.
endpoint, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil {
t.Fatalf("listen: %v", err)
}
defer func() {
_ = endpoint.Close()
}()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
"SWWAF_LOG_REMOTE_BUFFER": "1",
"SWWAF_METRICS_TOKEN": token,
}
out := runUntilStopped(t, env, func(url string) {
wantGreeting(t, url)
// More than one line has been written, and the buffer holds the
// last.
metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
for _, series := range []string{
"smallwebwaf_remote_log_lines_sent_total 0",
"smallwebwaf_remote_log_buffer_depth 1",
} {
if !strings.Contains(metrics, "\n"+series+"\n") {
t.Errorf("no %q in the metrics:\n%s", series, metrics)
}
}
if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") ||
!strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") {
t.Errorf("no line dropped in the metrics:\n%s", metrics)
}
// Closed, the endpoint refuses the connection made to send the
// lines still waiting at the stop, which then does not wait.
_ = endpoint.Close()
})
out.line(t, "type", "request")
}
func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
t.Parallel()
webhook := startWebhook(t)
rules := t.TempDir()
err := os.WriteFile(filepath.Join(rules, "50-app.rules"),
[]byte(`probe path ban ^/\.env$`+"\n"), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: rules,
"SWWAF_ALERT_WEBHOOK_URL": webhook.url,
"SWWAF_ALERT_WEBHOOK_HEADERS": "Authorization:Bearer " + adminSecret,
}
runUntilStopped(t, env, func(url string) {
// The probe bans the client, and the webhook is sent the alert.
wantRefused(t, url+".env")
post := webhook.waitFor(t, "ban", true)
if post.alert["client"] != localhost || post.alert["netblock"] != localhost+"/32" ||
post.authorization != "Bearer "+adminSecret {
t.Errorf("the webhook was sent %v, with Authorization %q", post.alert,
post.authorization)
}
// The webhook fails, so the alert for the ban made permanent by the
// client's next request waits.
webhook.failing.Store(true)
wantRefused(t, url)
webhook.waitFor(t, "permanent_ban", false)
})
// alerts.json keeps it as smallwebwaf stops, and once started again,
// smallwebwaf sends it.
var file struct {
Waiting []struct {
Event string `json:"event"`
} `json:"waiting"`
}
path := filepath.Join(dir, "alerts.json")
data, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
if err == nil {
err = json.Unmarshal(data, &file)
}
if err != nil || len(file.Waiting) != 1 || file.Waiting[0].Event != "permanent_ban" {
t.Fatalf("alerts.json holds %s (%v), want the permanent_ban alert waiting", data, err)
}
// It counts the alert sent in the metrics, read here from a client the
// ban does not cover.
const token = "0123456789abcdef0123456789abcdef"
webhook.failing.Store(false)
env["SWWAF_ALLOW_NETS"] = localhost
env["SWWAF_METRICS_TOKEN"] = token
runUntilStopped(t, env, func(url string) {
webhook.waitFor(t, "permanent_ban", true)
// As long as that takes, so that a slow test process cannot fail
// the test.
const sent = "\nsmallwebwaf_alerts_sent_total{destination=\"webhook\"} 1\n"
metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
for !strings.Contains(metrics, sent) {
time.Sleep(pollInterval)
metrics = metricsText(t, url+"_smallwebwaf/metrics", token)
}
for _, series := range []string{"failed", "suppressed", "dropped"} {
zero := "\nsmallwebwaf_alerts_" + series + "_total{destination=\"webhook\"} 0\n"
if !strings.Contains(metrics, zero) {
t.Errorf("no %q in the metrics:\n%s", zero, metrics)
}
}
})
}
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel()
@@ -349,7 +672,9 @@ func wantStartRefused(t *testing.T, dir, want string) {
out := &output{}
status := run(ctx, map[string]string{listenAddr: localhost + ":0", stateDir: dir}, out)
status := run(ctx, map[string]string{
listenAddr: localhost + ":0", stateDir: dir, rulesDir: t.TempDir(),
}, out)
if status != 1 {
t.Fatalf("exit status %d, want 1", status)
}
@@ -362,6 +687,30 @@ func wantStartRefused(t *testing.T, dir, want string) {
}
}
// wantRulesRefused runs smallwebwaf with its rule files in dir, and
// checks that it stops at start, with the error want. If it starts
// instead, it is stopped after waitLimit.
func wantRulesRefused(t *testing.T, dir, want string) {
t.Helper()
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
defer stop()
out := &output{}
status := run(ctx, map[string]string{
listenAddr: localhost + ":0", stateDir: t.TempDir(), rulesDir: dir,
}, out)
if status != 1 {
t.Fatalf("exit status %d, want 1", status)
}
line := out.line(t, "msg", "cannot use the rule files")
if line["error"] != want || line["level"] != "ERROR" {
t.Errorf("start refused with %v, want the error %q", line, want)
}
}
// startApp starts an app that answers every request with greeting, and
// returns its URL.
func startApp(t *testing.T) string {
@@ -436,14 +785,17 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
rateLimitPerDay: "50000",
"SWWAF_RATE_LIMIT_EXEMPT_PATHS": "",
"SWWAF_DENIED_COUNTRIES": "",
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
"SWWAF_BAN_RESPONSE": "403",
"SWWAF_LIMIT_BAN_DURATION": "1h",
"SWWAF_LIMIT_BAN_REPEAT_WINDOW": "24h",
"SWWAF_MAX_BAN_DURATION": "7d",
"SWWAF_ATTACK_BAN_DURATION": "7d",
"SWWAF_MAX_BANS": "5000",
"SWWAF_BAN_SCOPE_V4_PREFIX": "32",
"SWWAF_RULES_ENABLED": "true",
}
for name, value := range want {
@@ -483,6 +835,120 @@ 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)
}
// askAsAdmin sends a request with method to url, with body and
// adminSecret, and checks that it is answered 200.
func askAsAdmin(t *testing.T, method, url, body string) {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), method, url,
strings.NewReader(body))
if err != nil {
t.Fatalf("new request: %v", err)
}
req.Header.Set("Authorization", "Bearer "+adminSecret)
transport := &http.Transport{}
defer transport.CloseIdleConnections()
res, err := (&http.Client{Transport: transport}).Do(req)
if err != nil {
t.Fatalf("request: %v", err)
}
_ = res.Body.Close()
if res.StatusCode != http.StatusOK {
t.Fatalf("%s %s answered %d", method, url, res.StatusCode)
}
}
// onlyBan returns the one ban bans.json in dir holds.
func onlyBan(t *testing.T, dir string) map[string]any {
t.Helper()
path := filepath.Join(dir, "bans.json")
data, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
if err != nil {
t.Fatalf("read bans.json: %v", err)
}
var file struct {
Bans []map[string]any `json:"bans"`
}
err = json.Unmarshal(data, &file)
if err != nil || len(file.Bans) != 1 {
t.Fatalf("bans.json holds\n%s\nwant one ban (%v)", data, err)
}
return file.Bans[0]
}
// wantRefused checks that a request to url is refused with 403, the
// default SWWAF_BAN_RESPONSE.
func wantRefused(t *testing.T, url string) {
@@ -543,6 +1009,80 @@ func saveUntilAnswered(t *testing.T, path, content, url, from string, status int
}
}
// webhook is a stand-in for SWWAF_ALERT_WEBHOOK_URL. It notes each alert
// it is sent, and answers 204, or 503 while failing.
type webhook struct {
url string
failing atomic.Bool
mu sync.Mutex
posts []webhookPost
}
// webhookPost is an alert the webhook was sent, with the Authorization
// header sent with it, and whether the webhook took it.
type webhookPost struct {
alert map[string]any
authorization string
answered bool
}
// startWebhook starts a webhook that takes every alert.
func startWebhook(t *testing.T) *webhook {
t.Helper()
w := &webhook{}
server := httptest.NewServer(http.HandlerFunc(
func(rw http.ResponseWriter, r *http.Request) {
var alert map[string]any
_ = json.NewDecoder(r.Body).Decode(&alert)
failing := w.failing.Load()
w.mu.Lock()
w.posts = append(w.posts, webhookPost{
alert: alert, authorization: r.Header.Get("Authorization"),
answered: !failing,
})
w.mu.Unlock()
if failing {
rw.WriteHeader(http.StatusServiceUnavailable)
return
}
rw.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
w.url = server.URL + "/alerts"
return w
}
// waitFor waits until the webhook has been sent an alert for event that
// it took, or, unless answered, failed, and returns it. It waits as long
// as that takes, so that a slow test process cannot fail the test.
func (w *webhook) waitFor(t *testing.T, event string, answered bool) webhookPost {
t.Helper()
for {
w.mu.Lock()
for _, post := range w.posts {
if post.alert["event"] == event && post.answered == answered {
w.mu.Unlock()
return post
}
}
w.mu.Unlock()
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 {
+166 -41
View File
@@ -1,7 +1,8 @@
// 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,
// history, lookups.json GeoJS's answers, and alerts.json the cooldowns,
// the hour under way and the alerts waiting. 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
@@ -24,6 +25,7 @@ import (
"time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
@@ -42,12 +44,14 @@ const (
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
alertsJSON = "alerts.json"
)
var (
errVersion = errors.New("unknown version")
// errMissing is for an entry without a field it needs.
errMissing = errors.New("has no")
errCause = errors.New("is not limit, attack or admin")
)
// Params are what Load needs.
@@ -59,10 +63,13 @@ type Params struct {
// is (SWWAF_STATE_COUNTER_INTERVAL).
WriteDelay time.Duration
CounterInterval time.Duration
// Ledger, Limiter and GeoJS hold the state.
// Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
// receive a file_error alert for an edit set aside, and for a write
// that fails while smallwebwaf runs.
Ledger *bans.Ledger
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
Alerts *alerts.Queue
// Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC.
Now func() time.Time
@@ -91,15 +98,20 @@ type Files struct {
// bansFile is bans.json, indented for an admin to read and edit.
type bansFile struct {
Version int `json:"version"`
Bans []banEntry `json:"bans"`
Bans []BanEntry `json:"bans"`
}
// banEntry is a ban as bans.json holds it: a permanent ban's expires is
// null.
type banEntry struct {
// BanEntry is a ban as bans.json holds it: a permanent ban's expires is
// null, a ban an admin added may have no cause, which makes it an
// admin's, and lifted is left out until an admin lifts the ban. The ban
// endpoints answer with bans in this form too.
type BanEntry struct {
Netblock netip.Prefix `json:"netblock"`
Start time.Time `json:"start"`
Expires *time.Time `json:"expires"`
Cause string `json:"cause"`
Reason string `json:"reason,omitempty"`
Lifted *time.Time `json:"lifted,omitempty"`
Notes bans.Notes `json:"notes"`
}
@@ -115,6 +127,14 @@ type lookupsFile struct {
Lookups []lookup.Answer `json:"lookups"`
}
// alertsFile is alerts.json, indented for an admin to read and edit.
type alertsFile struct {
Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"`
Waiting []alerts.Alert `json:"waiting"`
}
// stateFile is the struct of a state file. Once the file is decoded, its
// check refuses the first entry without a field it needs, which would
// otherwise be read as something the entry does not say. data is the
@@ -141,23 +161,25 @@ func Load(params Params) (*Files, error) {
bansRead, bansErr := f.read(bansJSON)
clientsRead, clientsErr := f.read(clientsJSON)
lookupsRead, lookupsErr := f.read(lookupsJSON)
alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr)
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
if err != nil {
return nil, err
}
params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
"alerts_waiting", alertsRead)
return f, nil
}
// Run writes bans.json WriteDelay after a ban is made, with every ban
// made in between, and every file every CounterInterval, until ctx is
// done. A write that fails is logged, and the file is written again at
// its next write. Each write takes in an admin's edit of its file first,
// as writeFile describes.
// done. A write that fails is logged, raised as a file_error alert, and
// the file is written again at its next write. Each write takes in an
// admin's edit of its file first, as writeFile describes.
func (f *Files) Run(ctx context.Context) {
interval := time.NewTicker(f.params.CounterInterval)
defer interval.Stop()
@@ -175,9 +197,11 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue:
bansDue = nil
f.logFailure(f.writeFile(bansJSON))
f.logFailure(bansJSON, f.writeFile(bansJSON))
case <-interval.C:
f.logFailure(f.WriteAll())
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
f.logFailure(name, f.writeFile(name))
}
}
}
}
@@ -186,7 +210,7 @@ func (f *Files) Run(ctx context.Context) {
// fails does not keep the others from being written.
func (f *Files) WriteAll() error {
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON))
f.writeFile(lookupsJSON), f.writeFile(alertsJSON))
}
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
@@ -221,7 +245,7 @@ func (f *Files) Watch(ctx context.Context) {
return
case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON:
case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
f.fileChanged(name)
}
case err = <-watcher.Errors:
@@ -231,11 +255,22 @@ func (f *Files) Watch(ctx context.Context) {
}
}
// logFailure logs a write that failed.
func (f *Files) logFailure(err error) {
// logFailure logs a write of the state file name that failed, and raises
// a file_error alert for it.
func (f *Files) logFailure(name string, err error) {
if err != nil {
f.params.ProcessLog.Error("writing the state files failed",
"error", err.Error())
const failed = "writing the state files failed"
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: failed,
Detail: map[string]any{
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
},
})
f.params.ProcessLog.Error(failed, "error", err.Error())
}
}
@@ -260,7 +295,7 @@ func (f *Files) fileChanged(name string) {
// 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)
_, err := f.takeIn(name, data, true)
if err != nil {
return err
}
@@ -282,7 +317,7 @@ func (f *Files) read(name string) (int, error) {
return 0, err
}
return f.takeIn(name, data)
return f.takeIn(name, data, false)
}
// readChanged returns what the state file name holds, and whether that
@@ -306,9 +341,11 @@ func (f *Files) readChanged(name string) ([]byte, bool, error) {
// 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) {
// how many entries the file holds. edit is whether data is an admin's
// edit taken in while smallwebwaf runs, rather than the file read at the
// start. An error names the file and, where the JSON decoder tells it,
// the line and column, or else the entry.
func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
path := filepath.Join(f.params.Dir, name)
var entries int
@@ -327,7 +364,12 @@ func (f *Files) takeIn(name string, data []byte) (int, error) {
held = append(held, entry.ban())
}
f.params.Ledger.Load(held)
if edit {
f.params.Ledger.LoadEdit(held)
} else {
f.params.Ledger.Load(held)
}
entries = len(held)
case clientsJSON:
var file clientsFile
@@ -349,6 +391,18 @@ func (f *Files) takeIn(name string, data []byte) (int, error) {
f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups)
case alertsJSON:
var file alertsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Alerts.Load(alerts.State{
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
})
entries = len(file.Waiting)
}
f.sums[name] = sha256.Sum256(data)
@@ -400,8 +454,9 @@ func (f *Files) writeFile(name string) error {
// 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.
// the file the error is, and raises a file_error alert for it. 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)
@@ -410,8 +465,16 @@ func (f *Files) setAside(name string, parseErr error) error {
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())
const setAside = "set aside an edit of a state file that does not parse"
// Raised before it is logged, so that the alert is there once the log
// line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: setAside,
Detail: map[string]any{"file": path + ".bad", "error": parseErr.Error()},
})
f.params.ProcessLog.Error(setAside, "file", path+".bad", "error", parseErr.Error())
f.params.Metrics.StateFileEditSetAside(name)
return nil
@@ -422,12 +485,7 @@ func (f *Files) setAside(name string, parseErr error) error {
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))
}
file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
@@ -437,28 +495,66 @@ func (f *Files) encode(name string) ([]byte, error) {
return append(data, '\n'), nil
case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
default: // lookups.json
case lookupsJSON:
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
default: // alerts.json
held := f.params.Alerts.Snapshot()
file := alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting,
}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
}
}
// BanEntries returns held as bans.json lists them, an empty list for
// none.
func BanEntries(held []bans.Ban) []BanEntry {
entries := make([]BanEntry, 0, len(held))
for _, ban := range held {
entries = append(entries, newBanEntry(ban))
}
return entries
}
// 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}
func newBanEntry(ban bans.Ban) BanEntry {
entry := BanEntry{
Netblock: ban.Netblock, Start: ban.Start, Cause: ban.Cause, Reason: ban.Reason,
Notes: ban.Notes,
}
if !ban.Permanent() {
entry.Expires = &ban.Expires
}
if !ban.Lifted.IsZero() {
entry.Lifted = &ban.Lifted
}
return entry
}
// ban returns the ban an entry of bans.json holds.
func (e banEntry) ban() bans.Ban {
ban := bans.Ban{Netblock: e.Netblock, Start: e.Start, Notes: e.Notes}
func (e BanEntry) ban() bans.Ban {
ban := bans.Ban{
Netblock: e.Netblock, Start: e.Start, Cause: e.Cause, Reason: e.Reason,
Notes: e.Notes,
}
if e.Expires != nil {
ban.Expires = *e.Expires
}
if e.Lifted != nil {
ban.Lifted = *e.Lifted
}
return ban
}
@@ -466,7 +562,8 @@ func (e banEntry) ban() bans.Ban {
// client, a start, from which the length of the netblock's next ban is
// worked out, or an expires, which would make it permanent. A permanent
// ban's expires is null, which Bans cannot tell from a missing one, so
// each expires is read again as written.
// each expires is read again as written. A cause other than limit,
// attack or admin, most likely misspelt, is refused too.
func (f *bansFile) check(data []byte) error {
var written struct {
Bans []struct {
@@ -487,6 +584,9 @@ func (f *bansFile) check(data []byte) error {
return missing(i, "start")
case written.Bans[i].Expires == nil:
return missing(i, "expires")
case entry.Cause != "" && entry.Cause != bans.CauseLimit &&
entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin:
return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause)
}
}
@@ -544,6 +644,31 @@ func (f *lookupsFile) check(data []byte) error {
return nil
}
// check refuses a cooldown without its event or when its alert was sent,
// which would hold back no repeat, and an alert waiting without its event
// or its time.
func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns {
switch {
case cooldown.Event == "":
return fmt.Errorf("cooldowns %w", missing(i, "event"))
case cooldown.Sent.IsZero():
return fmt.Errorf("cooldowns %w", missing(i, "sent"))
}
}
for i, alert := range f.Waiting {
switch {
case alert.Event == "":
return fmt.Errorf("waiting %w", missing(i, "event"))
case alert.Time.IsZero():
return fmt.Errorf("waiting %w", missing(i, "time"))
}
}
return nil
}
// countsWithoutStart reports whether b holds requests but no start, which
// places them in time.
func countsWithoutStart(b ratelimit.Buckets) bool {
+470 -19
View File
@@ -10,8 +10,10 @@ import (
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"os"
"path/filepath"
"reflect"
"slices"
"strconv"
"strings"
@@ -19,6 +21,7 @@ import (
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
@@ -31,6 +34,7 @@ const (
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
alertsJSON = "alerts.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"
@@ -48,6 +52,8 @@ const permanentBansJSON = `{
"netblock": "2001:db8::/64",
"start": "2026-10-06T00:00:00Z",
"expires": null,
"cause": "admin",
"reason": "scrapes every commit",
"notes": {
"country": "DE",
"limit": 1000,
@@ -63,13 +69,87 @@ const permanentBansJSON = `{
},
"requests": 1500,
"refused": 3,
"earlier_bans": 5
"earlier_bans": {
"limit": 3,
"attack": 1,
"admin": 1
}
}
}
]
}
`
// liftedClient is the client whose ban liftedBansJSON holds.
const liftedClient = "203.0.113.9"
// liftedBansJSON is bans.json holding an hour's ban for a broken limit on
// liftedClient, from midnight, that an admin lifted ten minutes in.
const liftedBansJSON = `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", ` +
`"start": "2026-10-06T00:00:00Z", "expires": "2026-10-06T01:00:00Z", ` +
`"cause": "limit", "lifted": "2026-10-06T00:10:00Z"}]}`
// filledAlertsJSON is alerts.json holding the alerts of fill.
const filledAlertsJSON = `{
"version": 1,
"cooldowns": [
{
"event": "file_error",
"netblock": "",
"file": "/var/lib/smallwebwaf/bans.json",
"sent": "2026-10-06T00:00:00Z",
"suppressed_repeats": 0
},
{
"event": "ban",
"netblock": "203.0.113.9/32",
"sent": "2026-10-06T00:00:00Z",
"suppressed_repeats": 1
}
],
"hour": {
"start": "2026-10-06T00:00:00Z",
"sent": 2,
"held_back": {
"source_failure": 1
}
},
"waiting": [
{
"instance": "fsn1app1/gitea",
"time": "2026-10-06T00:00:00Z",
"event": "ban",
"client": "203.0.113.9",
"netblock": "203.0.113.9/32",
"asn": "",
"as_name": "",
"country": "DE",
"reason": "requests per minute over the limit of 1",
"detail": {
"cause": "limit"
},
"suppressed_repeats": 0
},
{
"instance": "fsn1app1/gitea",
"time": "2026-10-06T00:00:00Z",
"event": "file_error",
"client": "",
"netblock": "",
"asn": "",
"as_name": "",
"country": "",
"reason": "writing the state files failed",
"detail": {
"error": "no space left on device",
"file": "/var/lib/smallwebwaf/bans.json"
},
"suppressed_repeats": 0
}
]
}
`
func TestFilesWrittenAndReadBack(t *testing.T) {
t.Parallel()
@@ -96,13 +176,69 @@ func TestFilesWrittenAndReadBack(t *testing.T) {
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
}
// Each one-per-line file lists its entries by client, and nothing
// but the three files is left in the directory.
// but the four files is left in the directory.
wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
"192.0.2.1/32", "203.0.113.9/32")
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
}
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
got := readFile(t, filepath.Join(dir, alertsJSON))
if got != filledAlertsJSON {
t.Errorf("alerts.json\n%s\nwant\n%s", got, filledAlertsJSON)
}
}
func TestSourceFailureCooldownKeptInAlertsJSONAcrossARestart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
before := newParams(dir)
files := load(t, before)
failure := alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
Detail: map[string]any{"source": "geojs"},
}
before.Alerts.Raise(failure)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// After the restart, the cooldown read back holds back a repeat for the
// same source.
after := newParams(dir)
load(t, after)
after.Alerts.Raise(failure)
waiting := after.Alerts.Snapshot().Waiting
if len(waiting) != 1 || after.Alerts.Suppressed() != 1 {
t.Errorf("%d alerts wait and %d are held back, want the one read back and 1",
len(waiting), after.Alerts.Suppressed())
}
}
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
@@ -131,8 +267,10 @@ func TestMissingFilesAreEmptyState(t *testing.T) {
params := newParams(t.TempDir())
load(t, params)
held := params.Alerts.Snapshot()
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
len(params.GeoJS.Snapshot()) != 0 {
len(params.GeoJS.Snapshot()) != 0 || len(held.Cooldowns) != 0 ||
len(held.Waiting) != 0 || held.Hour.Sent != 0 {
t.Error("state from no files")
}
}
@@ -174,6 +312,11 @@ func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
`: netip.ParsePrefix("203.0.113.300/32")`,
},
{
"an unknown field of an alert waiting", alertsJSON,
`{"version": 1, "waiting": [{"event": "ban", "evnet": "ban"}]}`,
`: json: unknown field "evnet"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
@@ -263,10 +406,60 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
}
}
func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name, content string
// want is what the error says after the file's path.
want string
}{
{
"a cooldown without its event",
`{"version": 1, "cooldowns": [{"sent": "2026-10-06T00:00:00Z"}]}`,
`: cooldowns entry 1 has no "event"`,
},
{
"a cooldown without when it was sent",
`{"version": 1, "cooldowns": [{"event": "ban"}]}`,
`: cooldowns entry 1 has no "sent"`,
},
{
"an alert waiting without its event",
`{"version": 1, "waiting": [{"time": "2026-10-06T00:00:00Z"}]}`,
`: waiting entry 1 has no "event"`,
},
{
"an alert waiting without its time",
`{"version": 1, "waiting": [{"event": "ban"}]}`,
`: waiting entry 1 has no "time"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, alertsJSON, tc.content, tc.want)
})
}
}
func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
t.Parallel()
wantRefused(t, bansJSON, `{"version": 1, "bans": [`+
`{"netblock": "203.0.113.9/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": null, "cause": "attack"}, `+
`{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": null, "cause": "admin"}, `+
`{"netblock": "203.0.113.11/32", "start": "2026-10-06T00:00:00Z", `+
`"expires": null, "cause": "atack"}]}`,
`: entry 3's cause "atack" is not limit, attack or admin`)
}
func TestUnknownVersionStopsTheStart(t *testing.T) {
t.Parallel()
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} {
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
for _, content := range []string{`{"version": 2}`, `{}`} {
t.Run(file+" "+content, func(t *testing.T) {
t.Parallel()
@@ -316,12 +509,12 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
// A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it.
first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
first, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
midnight(), bans.Notes{})
time.Sleep(5 * time.Second)
second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
second, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
midnight(), bans.Notes{})
time.Sleep(5*time.Second - time.Nanosecond)
@@ -368,8 +561,8 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
time.Sleep(time.Nanosecond)
synctest.Wait()
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
}
})
}
@@ -412,6 +605,56 @@ func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) {
})
}
func TestWriteThatFailsWhileRunningRaisesAFileErrorAlertOncePerCooldown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.CounterInterval = time.Minute
run(t, load(t, params).Run)
// A directory in the way of clients.json's temporary file fails each
// of its writes, while the other files are written. It holds a file,
// so that the write cannot remove it.
err := os.Mkdir(filepath.Join(dir, clientsJSON+".tmp"), 0o700)
if err == nil {
err = os.WriteFile(filepath.Join(dir, clientsJSON+".tmp", "kept"), nil, 0o600)
}
if err != nil {
t.Fatalf("put a directory in the way: %v", err)
}
time.Sleep(time.Minute)
synctest.Wait()
waiting := params.Alerts.Snapshot().Waiting
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want 1", len(waiting))
}
message, _ := waiting[0].Detail["error"].(string)
if waiting[0].Event != alerts.EventFileError ||
waiting[0].Reason != "writing the state files failed" ||
waiting[0].Detail["file"] != filepath.Join(dir, clientsJSON) ||
!strings.Contains(message, clientsJSON+".tmp") {
t.Fatalf("alerts waiting %+v, want a file_error alert for clients.json, "+
"naming its temporary file", waiting)
}
// The next write fails too, within the cooldown, which holds it back.
time.Sleep(time.Minute)
synctest.Wait()
if len(params.Alerts.Snapshot().Waiting) != 1 || params.Alerts.Suppressed() != 1 {
t.Errorf("%d alerts wait and %d are held back, want 1 and 1",
len(params.Alerts.Snapshot().Waiting), params.Alerts.Suppressed())
}
})
}
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel()
@@ -537,7 +780,7 @@ func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
}
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
wantWriteFailed(t, params, bansJSON)
}
@@ -596,7 +839,7 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
`"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()}})
[]bans.Ban{{Netblock: client, Start: midnight(), Cause: bans.CauseAdmin}})
edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+
`{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`)
@@ -609,6 +852,25 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
wantTakenIn(t, lines, dir, lookupsJSON)
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
// A netblock with bits past its length is read as the netblock it is
// in.
edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+
`"netblock": "198.51.100.9/24", "sent": "2026-10-06T00:00:00Z"}], `+
`"waiting": [{"event": "file_error", "time": "2026-10-06T00:00:00Z"}]}`)
wantTakenIn(t, lines, dir, alertsJSON)
want := alerts.State{
Cooldowns: []alerts.Cooldown{{
Event: alerts.EventBan, Netblock: netip.MustParsePrefix("198.51.100.0/24"),
Sent: midnight(),
}},
Hour: alerts.Hour{HeldBack: map[string]int{}},
Waiting: []alerts.Alert{{Event: alerts.EventFileError, Time: midnight()}},
}
if got := params.Alerts.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("%s taken in as\n%+v\nwant\n%+v", alertsJSON, got, want)
}
}
func TestOwnWritesAreNotTakenIn(t *testing.T) {
@@ -680,7 +942,7 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned := params.Ledger.Check(client, midnight())
_, banned, _ := params.Ledger.Check(client, midnight())
if !banned {
t.Error("the ban added to bans.json does not refuse")
}
@@ -689,12 +951,106 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
wantTakenIn(t, lines, dir, bansJSON)
_, banned = params.Ledger.Check(client, midnight())
_, banned, _ = params.Ledger.Check(client, midnight())
if banned {
t.Error("the ban removed from bans.json still refuses")
}
}
func TestBanWithoutACauseTakenInAsAnAdmins(t *testing.T) {
t.Parallel()
const reason = "probes for logins"
// adminsBansJSON is bans.json as an admin writes it, with a ban on
// netblock without a cause, and adminsBans the bans it holds.
adminsBansJSON := func(netblock string) string {
return `{"version": 1, "bans": [{"netblock": "` + netblock + `", ` +
`"start": "2026-10-06T00:00:00Z", "expires": null, "reason": "` +
reason + `"}]}`
}
adminsBans := func(netblock string) []bans.Ban {
return []bans.Ban{{
Netblock: netip.MustParsePrefix(netblock),
Start: midnight(),
Cause: bans.CauseAdmin,
Reason: reason,
}}
}
// Read at the start, the ban is taken in as an admin's, though not
// counted among the bans made since the start, and written back with
// that cause and the admin's reason.
dir := t.TempDir()
edit(t, dir, bansJSON, adminsBansJSON("203.0.113.0/24"))
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
wantEqual(t, bansJSON, params.Ledger.Snapshot(), adminsBans("203.0.113.0/24"))
wantMadeByAnAdmin(t, params.Ledger, 0)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
var written struct {
Bans []struct {
Cause string `json:"cause"`
Reason string `json:"reason"`
} `json:"bans"`
}
err = json.Unmarshal([]byte(readFile(t, filepath.Join(dir, bansJSON))), &written)
if err != nil || len(written.Bans) != 1 || written.Bans[0].Cause != bans.CauseAdmin ||
written.Bans[0].Reason != reason {
t.Errorf("bans.json holds %+v (%v), want the ban with the cause admin "+
"and the reason %q", written, err, reason)
}
// Taken in while smallwebwaf runs, a ban on another netblock is an
// admin's too, and one made since the start.
watch(t, files, lines)
edit(t, dir, bansJSON, adminsBansJSON("198.51.100.0/24"))
wantTakenIn(t, lines, dir, bansJSON)
wantEqual(t, bansJSON, params.Ledger.Snapshot(), adminsBans("198.51.100.0/24"))
wantMadeByAnAdmin(t, params.Ledger, 1)
}
func TestLiftedBanReadAtTheStart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
edit(t, dir, bansJSON, liftedBansJSON)
params := newParams(dir)
wantLiftedBanKept(t, load(t, params), dir, params.Ledger)
}
func TestBanLiftedByAnEditWhileRunning(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
files := load(t, params)
watch(t, files, lines)
// The ban that liftedBansJSON lifts, before it is lifted.
netblock := netip.MustParsePrefix(liftedClient + "/32")
params.Ledger.BanForLimit(netblock, midnight(), bans.Notes{})
_, banned, _ := params.Ledger.Find(netblock.Addr(), afterLifting())
if !banned {
t.Fatal("the ban does not refuse before it is lifted")
}
edit(t, dir, bansJSON, liftedBansJSON)
wantTakenIn(t, lines, dir, bansJSON)
wantLiftedBanKept(t, files, dir, params.Ledger)
}
func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Parallel()
@@ -722,7 +1078,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
edit(t, dir, bansJSON, broken)
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, clientsJSON)
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantFiles(t, dir, alertsJSON, 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.
@@ -739,7 +1095,14 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Errorf("set aside with %v", line)
}
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
// It is raised as a file_error alert, with the same file and error.
waiting := params.Alerts.Snapshot().Waiting
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
waiting[0].Detail["file"] != path+".bad" || waiting[0].Detail["error"] != message {
t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad")
}
wantFiles(t, dir, alertsJSON, 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)
@@ -868,7 +1231,8 @@ func midnight() time.Time {
}
// newParams returns Params for the state files in dir, with parts that
// hold nothing yet. GeoJS is never asked.
// hold nothing yet. GeoJS is never asked, and the alerts, at most two an
// hour, are never sent.
func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler)
m := metrics.New(1)
@@ -881,26 +1245,39 @@ func newParams(dir string) state.Params {
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
AttackBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
}),
Limiter: ratelimit.New(ratelimit.Limits{}),
GeoJS: lookup.New(lookup.Params{
Now: midnight, ProcessLog: discard, Metrics: m,
}),
Alerts: alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 2,
Instance: "fsn1app1/gitea",
Now: midnight,
}),
Now: midnight,
ProcessLog: discard,
Metrics: m,
}
}
// fill puts a ban that ends and one that does not, clients with counts
// and histories, and GeoJS answers into the parts of params.
// fill puts a permanent ban an admin made, a ban for a broken limit and
// one for a clear sign of attack, clients with counts and histories,
// GeoJS answers, and alerts, as filledAlertsJSON holds them, into the
// parts of params.
func fill(params state.Params) {
now := midnight()
client := netip.MustParsePrefix("203.0.113.9/32")
params.Ledger.Load([]bans.Ban{permanentBan()})
params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1})
params.Ledger.BanForAttack(netip.MustParsePrefix("192.0.2.1/32"), now,
bans.Notes{RuleID: "env-file", Target: "path"})
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
params.Limiter.Count(netip.MustParsePrefix(c), now)
@@ -917,6 +1294,25 @@ func fill(params state.Params) {
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
},
})
// An alert waiting, a repeat of it the cooldown holds back, another
// alert waiting, and one past the two an hour, for the hour's summary.
ban := alerts.Alert{
Event: alerts.EventBan, Client: client.Addr(), Netblock: client, Country: "DE",
Reason: "requests per minute over the limit of 1",
Detail: map[string]any{"cause": "limit"},
}
params.Alerts.Raise(ban)
params.Alerts.Raise(ban)
params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError, Reason: "writing the state files failed",
Detail: map[string]any{
"file": "/var/lib/smallwebwaf/bans.json", "error": "no space left on device",
},
})
params.Alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
})
}
// permanentBan is the ban permanentBansJSON holds.
@@ -924,6 +1320,8 @@ func permanentBan() bans.Ban {
return bans.Ban{
Netblock: netip.MustParsePrefix("2001:db8::/64"),
Start: midnight(),
Cause: bans.CauseAdmin,
Reason: "scrapes every commit",
Notes: bans.Notes{
Country: "DE",
Limit: 1000,
@@ -939,11 +1337,64 @@ func permanentBan() bans.Ban {
},
Requests: 1500,
Refused: 3,
EarlierBans: 5,
EarlierBans: bans.EarlierBans{Limit: 3, Attack: 1, Admin: 1},
},
}
}
// afterLifting is a time after the ban liftedBansJSON holds was lifted,
// while it would still last.
func afterLifting() time.Time {
return midnight().Add(30 * time.Minute)
}
// wantLiftedBanKept checks that ledger holds the ban liftedBansJSON holds,
// which refuses nothing and does not make the next ban for a broken limit
// longer, and that files write it to bans.json, in dir, still lifted.
func wantLiftedBanKept(
t *testing.T, files *state.Files, dir string, ledger *bans.Ledger,
) {
t.Helper()
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
const lifted = `"lifted": "2026-10-06T00:10:00Z"`
if got := readFile(t, filepath.Join(dir, bansJSON)); !strings.Contains(got, lifted) {
t.Errorf("bans.json holds\n%s\nwant the ban with %s", got, lifted)
}
netblock := netip.MustParsePrefix(liftedClient + "/32")
_, banned, _ := ledger.Check(netblock.Addr(), afterLifting())
if banned {
t.Error("the lifted ban refuses")
}
// Were the lifted ban counted, the next would last three hours.
ban, _ := ledger.BanForLimit(netblock, afterLifting(), bans.Notes{})
if ban.Expires.Sub(ban.Start) != time.Hour {
t.Errorf("the next ban lasts %s, want 1h", ban.Expires.Sub(ban.Start))
}
held := ledger.Bans(netblock)
if len(held) != 2 || !held[0].Lifted.Equal(midnight().Add(10*time.Minute)) {
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
}
}
// wantMadeByAnAdmin checks how many bans ledger counts as made by an
// admin since the start.
func wantMadeByAnAdmin(t *testing.T, ledger *bans.Ledger, want int) {
t.Helper()
if got := ledger.Made(bans.CauseAdmin); got != want {
t.Errorf("%d bans made by an admin, want %d", got, want)
}
}
// load reads the state files into the parts of params.
func load(t *testing.T, params state.Params) *state.Files {
t.Helper()
+27 -5
View File
@@ -3,9 +3,10 @@
# deploy/example-app, then run the app's container with a volume for the
# state files and check that the health check passes, that a request is
# served through smallwebwaf, that a second one in a minute bans the
# client, that `sv stop` stops smallwebwaf in order, that `docker stop`
# stops the container without having to kill it, and that a new
# container on the same volume still refuses the banned client. The
# client, that a probe for /.env bans another client, which its next
# request bans for good, that `sv stop` stops smallwebwaf in order, that
# `docker stop` stops the container without having to kill it, and that
# a new container on the same volume still refuses the banned client. The
# containers, the volume and both images are removed however the script
# ends. Building the app needs network access, for nixpkgs' binary cache.
# script/check does not run this.
@@ -52,9 +53,13 @@ healthy() {
[ "$status" = healthy ]
}
# logged <text>: the container's output holds text.
# logged <text>...: a line of the container's output holds every text,
# in any order.
logged() {
docker logs "$CONTAINER" 2>&1 | grep -qF "$1"
lines="$(docker logs "$CONTAINER" 2>&1)"
for text in "$@"; do
lines="$(printf '%s\n' "$lines" | grep -F "$text")" || return 1
done
}
# start_container: run the app's container, with the state files on the
@@ -77,6 +82,15 @@ refused() {
[ "$code" = 403 ]
}
# refused_from <client> <path>: a request for path from client, as
# X-Forwarded-For names it, gets 403. smallwebwaf believes the header
# from docker's gateway, a private address.
refused_from() {
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
--max-time 10 --header "X-Forwarded-For: $1" "http://$address$2")" || true
[ "$code" = 403 ]
}
main() {
cd "$ROOT"
trap cleanup EXIT
@@ -100,6 +114,14 @@ main() {
wait_for "smallwebwaf logged no ban" logged '"action":"rate_limited"'
echo "example-app: a second request in a minute bans the client"
refused_from 203.0.113.9 /.env || fail "a probe for /.env was not refused"
wait_for "smallwebwaf logged no ban for the probe" \
logged '"action":"banned"' '"rule_ids":["env-file"]'
refused_from 203.0.113.9 / || fail "the client of the probe was let through"
wait_for "the client's next request did not make its ban permanent" \
logged '"ban_expires":"permanent"'
echo "example-app: a probe for /.env bans the client, its next request for good"
docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null ||
fail "sv stop smallwebwaf failed"
wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"'
+7 -1
View File
@@ -1,7 +1,9 @@
#!/bin/sh
# script/run: build bin/smallwebwaf with script/build and run it, with
# the settings in the environment. Unless SWWAF_STATE_DIR is set, the
# state files go in bin/state, beside the binary.
# state files go in bin/state, beside the binary, and unless
# SWWAF_RULES_DIR is set, the rule files are those of share/rules.d,
# which the image ships.
set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -14,6 +16,10 @@ main() {
export SWWAF_STATE_DIR
mkdir -p "$SWWAF_STATE_DIR"
fi
if [ -z "${SWWAF_RULES_DIR+set}" ]; then
SWWAF_RULES_DIR="$ROOT/share/rules.d"
export SWWAF_RULES_DIR
fi
exec "$ROOT/bin/smallwebwaf"
}
+15
View File
@@ -0,0 +1,15 @@
# 00-default.rules: probes no real visitor sends, anchored at the site root
# id target action regex
env-file path ban (?i)^/\.env(\.[a-z]+)?$
vcs-dir path ban (?i)^/\.(git|svn|hg|bzr)(/|$)
secrets-dir path ban (?i)^/\.(aws|ssh|docker|kube)/
secret-file path ban (?i)^/\.(htpasswd|htaccess|npmrc|netrc|pgpass|git-credentials|bash_history|DS_Store)$
editor-dir path ban (?i)^/\.(vscode|idea)/
backup-file path ban (?i)^/[^/]+\.(php(\.[a-z0-9]+|~)|sql(\.[a-z0-9]+)?)$
log-file path ban (?i)^/(debug|error|access)\.log$
compose-file path ban (?i)^/(docker-)?compose\.ya?ml$
php-shell path ban (?i)^/(shell|c99|r57|wso|alfa)\.php$
scanner-agent user_agent ban (?i)\b(sqlmap|nikto|nuclei|masscan|zgrab|wpscan)\b
path-traversal uri block (\.\./){2,}
empty-agent user_agent log ^$