Serve Prometheus metrics behind SWWAF_METRICS_TOKEN (closes #23)
check / check (push) Successful in 4m12s

GET /_smallwebwaf/metrics answers in the Prometheus text format for a
request carrying SWWAF_METRICS_TOKEN, 401 without it and 404 while it is
unset. Every request under /_smallwebwaf/ but the health check now goes
through the checks and is answered where it would be forwarded, 404 for
any path but the metrics, so none reaches the app. SWWAF_METRICS_TOP_N
bounds the series by country, the rest counted as other.

Judgement call: a request answered at smallwebwaf's own endpoints is
neither forwarded nor refused in the client's history.
Deviation: go.mod and go.sum written by hand from the module proxy and
sum.golang.org, as go runs only through make.
Deviation: no metrics yet for state files read again after an edit or
edits set aside; that work is not merged.

Model: opus-5-5
This commit is contained in:
2026-10-06 08:27:11 +00:00
parent 68f687cb0c
commit 2776bb4b09
23 changed files with 1459 additions and 66 deletions
+35
View File
@@ -112,6 +112,8 @@ type Ledger struct {
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans.
held int
// made is how many bans BanForLimit has made since the start.
made 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
@@ -206,6 +208,7 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
Notes: notes,
}
l.add(ban)
l.made++
select {
case l.changed <- struct{}{}:
@@ -229,6 +232,38 @@ 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 {
l.mu.Lock()
defer l.mu.Unlock()
return l.made
}
// Count returns how many of the bans held are active at now, and how many
// are permanent.
func (l *Ledger) Count(now time.Time) (int, int) {
l.mu.Lock()
defer l.mu.Unlock()
active, permanent := 0, 0
for _, bans := range l.netblocks.Values() {
for _, ban := range *bans {
if ban.ActiveAt(now) {
active++
}
if ban.Permanent() {
permanent++
}
}
}
return active, permanent
}
// Snapshot returns every ban held, sorted by netblock, and each
// netblock's bans oldest first, as bans.json lists them.
func (l *Ledger) Snapshot() []Ban {
+34
View File
@@ -17,6 +17,7 @@ import (
"strconv"
"strings"
"time"
"unicode/utf8"
)
// Config is smallwebwaf's settings. A timeout, size or rate limit of zero
@@ -103,6 +104,12 @@ type Config struct {
StateDir string
StateWriteDelay time.Duration
StateCounterInterval time.Duration
// 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
// settings are the values read, as given or by default, for the
// log line at start.
@@ -119,6 +126,10 @@ const (
mebibyte = 1 << 20
gibibyte = 1 << 30
ipv4Bits = 32
// minTokenLength is the fewest characters a token may have.
minTokenLength = 32
// masked is what the log shows for a token that is set.
masked = "********"
)
var (
@@ -150,6 +161,7 @@ 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")
)
// FromEnvironment reads the settings with lookupEnv, normally
@@ -188,6 +200,8 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
}
for _, country := range cfg.ExclusivelyAllowedCountries {
@@ -353,6 +367,26 @@ func (e *environment) absolutePath(name, defaultValue string) string {
return path
}
// token reads a setting that is a bearer token. Unset, it is "", which
// switches off what it guards; set, it must be at least minTokenLength
// characters. Neither the log nor an error shows its value.
func (e *environment) token(name string) string {
value, set := e.lookupEnv(name)
if !set {
e.settings = append(e.settings, slog.String(name, ""))
return ""
}
e.settings = append(e.settings, slog.String(name, masked))
if utf8.RuneCountInString(value) < minTokenLength {
e.check(name, errShortToken)
}
return value
}
// 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) {
+54 -1
View File
@@ -44,8 +44,13 @@ const (
stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
)
// token is a token of 32 characters, the shortest allowed.
const token = "0123456789abcdef0123456789abcdef"
// off switches a timeout, a size limit or a rate limit off.
const off = "off"
@@ -98,6 +103,8 @@ func TestDefaults(t *testing.T) {
StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute,
MetricsToken: "",
MetricsTopN: 50,
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -145,6 +152,8 @@ func TestValuesAsSet(t *testing.T) {
stateDir: "/srv/waf-state",
stateWriteDelay: "500ms",
stateCounterInterval: "1h",
metricsToken: token,
metricsTopN: "10",
})
wantSettings(t, cfg, config.Config{
@@ -169,6 +178,8 @@ func TestValuesAsSet(t *testing.T) {
StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour,
MetricsToken: token,
MetricsTopN: 10,
})
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -344,6 +355,7 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
@@ -360,6 +372,39 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
}
}
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
// Characters are counted, not bytes: each é takes two.
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv)
want := metricsToken + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestTokenIsLoggedMasked(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{metricsToken: token})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
if strings.Contains(out.String(), token) ||
!strings.Contains(out.String(), `"`+metricsToken+`":"********"`) {
t.Errorf("the token is not logged masked: %s", out.String())
}
}
func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Parallel()
@@ -407,6 +452,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
metricsToken: "",
metricsTopN: "50",
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -435,7 +482,8 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans and the state files.
// wantBanSettings checks the settings for bans, the state files and the
// metrics.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
@@ -453,6 +501,11 @@ func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
got.StateCounterInterval != want.StateCounterInterval {
t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
}
if got.MetricsToken != want.MetricsToken || got.MetricsTopN != want.MetricsTopN {
t.Errorf("metrics token %q and top %d, want %q and %d",
got.MetricsToken, got.MetricsTopN, want.MetricsToken, want.MetricsTopN)
}
}
// wantNetblocks checks a list of netblocks.
+17
View File
@@ -19,6 +19,7 @@ import (
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/metrics"
)
// URL is GeoJS's country endpoint. Asked about several addresses at once,
@@ -64,6 +65,9 @@ type Params struct {
Now func() time.Time
// ProcessLog receives GeoJS's failures.
ProcessLog *slog.Logger
// Metrics count the requests to GeoJS, those that failed, and the
// clients that go without an answer.
Metrics *metrics.Metrics
}
// GeoJS looks up clients' countries through GeoJS. At most one request
@@ -73,6 +77,7 @@ type GeoJS struct {
url string
now func() time.Time
processLog *slog.Logger
metrics *metrics.Metrics
// httpClient follows no redirect, so that visitors' addresses go to
// GeoJS alone: a redirect is a failure.
httpClient *http.Client
@@ -121,6 +126,7 @@ func New(params Params) *GeoJS {
url: params.URL,
now: params.Now,
processLog: params.ProcessLog,
metrics: params.Metrics,
httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
@@ -160,6 +166,9 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
defer g.mu.Unlock()
country, found := g.kept(client)
if !found {
g.metrics.GeoJSUnanswered.Inc()
}
w, waiting := g.waiting[client]
if !found && waiting {
@@ -234,6 +243,8 @@ func (g *GeoJS) answerOrWait(
g.ask(ctx)
if w == nil {
g.metrics.GeoJSUnanswered.Inc()
return "", nil // too many clients wait already
}
@@ -243,6 +254,8 @@ func (g *GeoJS) answerOrWait(
}
if w.late {
g.metrics.GeoJSUnanswered.Inc()
return "", nil
}
@@ -355,6 +368,8 @@ func (g *GeoJS) keep(
}
if err != nil {
g.metrics.GeoJSFailures.Inc()
g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay),
maxRetryDelay)
g.retryAt = now.Add(g.retryDelay)
@@ -399,6 +414,8 @@ func (g *GeoJS) request(
req.URL.RawQuery = "ip=" + strings.Join(addrs, ",")
g.metrics.GeoJSRequests.Inc()
res, err := g.httpClient.Do(req)
if err != nil {
// Do's error names the URL, and so the visitors' addresses, which
+3
View File
@@ -14,6 +14,7 @@ import (
"time"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
)
const (
@@ -194,6 +195,7 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
URL: lookup.URL,
Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1),
})
g.SetTransport(geojs)
@@ -493,6 +495,7 @@ func start() (*standIn, *testClock, *lookup.GeoJS) {
URL: lookup.URL,
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1),
})
g.SetTransport(geojs)
+116
View File
@@ -0,0 +1,116 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label under which the countries outside the busiest are
// counted.
const other = "other"
// countries are the metrics by the client's country, for requests whose
// client's country is known. The topN busiest countries, by their requests
// since the start, have series of their own, and the others are counted
// under other, so that there are never more than topN + 1 series. A
// country that drops out of the busiest loses its series, and its next
// requests are counted under other; one that becomes one of them gets a
// series that counts from then on. Each series therefore only ever goes
// up.
type countries struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are the requests the country lists refused.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's requests since the start, by which the
// countries are ranked. GeoJS gives two-letter codes, so it holds at
// most a few hundred.
seen map[string]int64
// top are the countries with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *countries {
byCountry := []string{"country"}
return &countries{
topN: topN,
requests: counterVec("smallwebwaf_country_requests_total",
"Requests, by the client's country.", byCountry),
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
"Request body bytes, by the client's country.", byCountry),
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
"Response body bytes, by the client's country.", byCountry),
refused: counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
byCountry),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// add counts a request from its log line, whose country is known.
func (c *countries) add(line *requestlog.Line) {
c.mu.Lock()
defer c.mu.Unlock()
c.seen[line.Country]++
label := c.label(line.Country)
c.requests.WithLabelValues(label).Inc()
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if line.Action == requestlog.ActionCountryDenied {
c.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from country is counted under: the
// country while it is one of the busiest, other while it is not. A
// country busier than the least busy of them takes its place, and that
// country's series are dropped.
func (c *countries) label(country string) string {
if c.top[country] {
return country
}
if len(c.top) < c.topN {
c.top[country] = true
return country
}
least := ""
for top := range c.top {
if least == "" || c.seen[top] < c.seen[least] {
least = top
}
}
if c.seen[country] <= c.seen[least] {
return other
}
delete(c.top, least)
for _, vec := range []*prometheus.CounterVec{
c.requests, c.requestBytes, c.responseBytes, c.refused,
} {
vec.DeleteLabelValues(least)
}
c.top[country] = true
return country
}
+257
View File
@@ -0,0 +1,257 @@
// Package metrics keeps smallwebwaf's Prometheus metrics, as the "Metrics
// endpoint" section of SPEC.md lists them, and serves them in the
// Prometheus text format. No metric carries a client's address.
package metrics
import (
"net/http"
"strconv"
"time"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/collectors"
"github.com/prometheus/client_golang/prometheus/promhttp"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
type Metrics struct {
registry *prometheus.Registry
handler http.Handler
inFlight prometheus.Gauge
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
requestDuration prometheus.Histogram
upstreamDuration prometheus.Histogram
rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec
countries *countries
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
// that failed. GeoJSUnanswered are the requests whose client counted
// as coming from an unknown country because GeoJS had not answered
// about it in time.
GeoJSRequests prometheus.Counter
GeoJSFailures prometheus.Counter
GeoJSUnanswered prometheus.Counter
stateFileWrites *prometheus.CounterVec
stateFileWriteFailures *prometheus.CounterVec
stateFileLastWrite *prometheus.GaugeVec
stateFileSize *prometheus.GaugeVec
}
// New returns the metrics, with the Go runtime's and the process's own.
// topN is how many countries get series of their own
// (SWWAF_METRICS_TOP_N).
func New(topN int) *Metrics {
byStatus := []string{"status_class", "action"}
byFile := []string{"file"}
m := &Metrics{
registry: prometheus.NewRegistry(),
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
Name: "smallwebwaf_requests_in_flight",
Help: "Requests under way.",
}),
requests: counterVec("smallwebwaf_requests_total",
"Requests, by the class of their status and their action.", byStatus),
requestBytes: counterVec("smallwebwaf_request_bytes_total",
"Request body bytes, by the class of the status and the action.",
byStatus),
responseBytes: counterVec("smallwebwaf_response_bytes_total",
"Response body bytes, by the class of the status and the action.",
byStatus),
requestDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
Name: "smallwebwaf_request_duration_seconds",
Help: "How long requests took, from their arrival to their end.",
}),
upstreamDuration: prometheus.NewHistogram(prometheus.HistogramOpts{
Name: "smallwebwaf_upstream_duration_seconds",
Help: "How long requests passed to the app took, from then to their end.",
}),
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
"Requests that broke a rate limit, by its window.",
[]string{"window"}),
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
"Requests that passed a size or time limit, by its setting.",
[]string{"limit"}),
offences: counterVec("smallwebwaf_offences_total",
"Offences, by kind.", []string{"kind"}),
countries: newCountries(topN),
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_requests_total",
Help: "Requests to GeoJS.",
}),
GeoJSFailures: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_failures_total",
Help: "Requests to GeoJS that failed.",
}),
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_unanswered_total",
Help: "Requests whose client counted as coming from an unknown " +
"country because GeoJS had not answered about it in time.",
}),
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
"Writes of each state file.", byFile),
stateFileWriteFailures: counterVec("smallwebwaf_state_file_write_failures_total",
"Writes of each state file that failed.", byFile),
stateFileLastWrite: gaugeVec("smallwebwaf_state_file_last_write_timestamp_seconds",
"When each state file was last written, in seconds since 1970.", byFile),
stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes",
"The size of each state file, as it was last written.", byFile),
}
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
m.registry.MustRegister(
collectors.NewGoCollector(),
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
m.inFlight, m.requests, m.requestBytes, m.responseBytes,
m.requestDuration, m.upstreamDuration,
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
m.countries.refused,
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize,
)
return m
}
// AddBansAndClients adds the metrics read from the ledger and the table
// of clients as the metrics are asked for: the bans made since the start,
// 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{
Name: "smallwebwaf_bans_made_total",
Help: "Bans made, by cause.",
ConstLabels: prometheus.Labels{"cause": "limit"},
}, func() float64 {
return float64(ledger.Made())
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_active_bans",
Help: "Bans active now, the permanent ones included.",
}, func() float64 {
active, _ := ledger.Count(now())
return float64(active)
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_permanent_bans",
Help: "Permanent bans.",
}, func() float64 {
_, permanent := ledger.Count(now())
return float64(permanent)
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_tracked_clients",
Help: "Clients in the table of clients.",
}, func() float64 {
return float64(limiter.Len())
}),
)
}
// ServeHTTP answers with the metrics in the Prometheus text format.
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
m.handler.ServeHTTP(w, r)
}
// RequestStarted counts a request as under way.
func (m *Metrics) RequestStarted() {
m.inFlight.Inc()
}
// RequestEnded counts a request that has ended, from its log line. limit
// is the setting whose size or time limit the request passed, "" if none.
// duration is how long the request took, and upstreamDuration how long it
// took from when it was passed to the app, zero if it was not.
func (m *Metrics) RequestEnded(
line *requestlog.Line, limit string, duration, upstreamDuration time.Duration,
) {
m.inFlight.Dec()
class := statusClass(line.Status)
m.requests.WithLabelValues(class, line.Action).Inc()
m.requestBytes.WithLabelValues(class, line.Action).Add(float64(line.RequestBytes))
m.responseBytes.WithLabelValues(class, line.Action).Add(float64(line.ResponseBytes))
m.requestDuration.Observe(duration.Seconds())
if upstreamDuration > 0 {
m.upstreamDuration.Observe(upstreamDuration.Seconds())
}
if line.LimitHit != "" {
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
}
if limit != "" {
m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc()
}
if line.Offence != "" {
m.offences.WithLabelValues(line.Offence).Inc()
}
if line.Country != "" {
m.countries.add(line)
}
}
// StateFileWritten counts a write of the state file name, of size bytes,
// that ended with err.
func (m *Metrics) StateFileWritten(name string, size int, err error) {
m.stateFileWrites.WithLabelValues(name).Inc()
// The series of failures is there from the first write, at zero until
// one fails.
failures := m.stateFileWriteFailures.WithLabelValues(name)
if err != nil {
failures.Inc()
return
}
m.stateFileLastWrite.WithLabelValues(name).SetToCurrentTime()
m.stateFileSize.WithLabelValues(name).Set(float64(size))
}
// statusClass returns the class of status, such as 2xx, or none when no
// status was sent.
func statusClass(status int) string {
if status == 0 {
return "none"
}
// A status's class is its hundreds: 404 is in 4xx.
const hundred = 100
return strconv.Itoa(status/hundred) + "xx"
}
// counterVec returns a counter named name, described by help, with a
// series for each set of values of labels.
func counterVec(name, help string, labels []string) *prometheus.CounterVec {
return prometheus.NewCounterVec(prometheus.CounterOpts{Name: name, Help: help},
labels)
}
// gaugeVec returns a gauge named name, described by help, with a series
// for each set of values of labels.
func gaugeVec(name, help string, labels []string) *prometheus.GaugeVec {
return prometheus.NewGaugeVec(prometheus.GaugeOpts{Name: name, Help: help}, labels)
}
+40
View File
@@ -0,0 +1,40 @@
package proxy
import (
"crypto/subtle"
"net/http"
"strings"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// answerAdmin answers a request for smallwebwaf itself, under
// /_smallwebwaf/, once it has passed the checks: GET MetricsPath with
// SWWAF_METRICS_TOKEN gets the metrics, and without it 401. Any other
// request gets 404, as the metrics do while SWWAF_METRICS_TOKEN is unset.
func (rq *request) answerAdmin() {
rq.line.Action = requestlog.ActionAdmin
rq.startClientResponseTimeout()
token := rq.h.config.MetricsToken
switch {
case token == "" || rq.in.Method != http.MethodGet || rq.in.URL.Path != MetricsPath:
http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound)
case !hasToken(rq.in, token):
rq.out.Header().Set("WWW-Authenticate", "Bearer")
http.Error(rq.out, http.StatusText(http.StatusUnauthorized),
http.StatusUnauthorized)
default:
rq.h.metrics.ServeHTTP(rq.out, rq.in)
}
}
// hasToken reports whether r carries token, as Authorization: Bearer
// <token>.
func hasToken(r *http.Request, token string) bool {
scheme, sent, _ := strings.Cut(r.Header.Get("Authorization"), " ")
return strings.EqualFold(scheme, "Bearer") &&
subtle.ConstantTimeCompare([]byte(sent), []byte(token)) == 1
}
+25 -7
View File
@@ -402,36 +402,54 @@ func (s *sender) get(from string, status int, action string) logLine {
func (s *sender) request(from, path string, status int, action string) logLine {
s.t.Helper()
line, _ := s.requestWithHeader(from, path, "", status, action)
return line
}
// requestWithHeader is request with header, such as "Authorization:
// Bearer x", added to the request unless it is "". It returns the body of
// the answer too.
func (s *sender) requestWithHeader(
from, path, header string, status int, action string,
) (logLine, string) {
s.t.Helper()
if header != "" {
header += "\r\n"
}
conn := dial(s.t, s.addr)
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n\r\n")
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+
header+"\r\n")
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
s.t.Fatalf("set read deadline: %v", err)
}
got := 0
var got answer
res, err := http.ReadResponse(bufio.NewReader(conn), nil)
switch {
case err == nil:
got = readAnswer(res).status
got = readAnswer(res)
case !errors.Is(err, io.ErrUnexpectedEOF):
s.t.Fatalf("read response: %v", err)
}
_ = conn.Close()
if got != status {
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, got,
status)
if got.status != status {
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from,
got.status, status)
}
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
s.sent++
wantLine(s.t, line, status, action)
return line
return line, string(got.body)
}
+2
View File
@@ -46,6 +46,7 @@ func (b *requestBody) Read(p []byte) (int, error) {
b.rq.refuse(refusal{
status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge,
limit: "SWWAF_REQUEST_MAX_BYTES",
})
}
@@ -81,6 +82,7 @@ func (b *responseBody) Read(p []byte) (int, error) {
b.rq.refuse(refusal{
status: http.StatusBadGateway,
action: requestlog.ActionTooLarge,
limit: "SWWAF_RESPONSE_MAX_BYTES",
})
return n, errResponseTooLarge
+18
View File
@@ -86,6 +86,24 @@ func TestHealthEndpointIsNotInTheHistory(t *testing.T) {
}
}
func TestRequestForSmallwebwafIsNeitherForwardedNorRefused(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now,
map[string]string{metricsToken: token})
scrape(t, addr)
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
out.requestLines(t, 2)
history := historyOf(t, server, localhost)
if history.Requests != 2 || history.Forwarded != 0 || history.Refused != 0 {
t.Errorf("history counts %d requests, %d forwarded and %d refused, "+
"want 2, 0 and 0", history.Requests, history.Forwarded, history.Refused)
}
}
// historyOf returns the history of the client at addr.
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
t.Helper()
+466
View File
@@ -0,0 +1,466 @@
package proxy_test
import (
"bytes"
"io"
"net/http"
"net/http/httptest"
"net/netip"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
// token is the SWWAF_METRICS_TOKEN the tests set, and bearer how a
// request carries it.
token = "0123456789abcdef0123456789abcdef"
bearer = "Bearer " + token
)
func TestMetricsAreOffWhileTheTokenIsUnset(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, out := startProxy(t, app.URL, nil)
// An empty token does not match the unset one either.
for i, authorization := range []string{bearer, "Bearer ", ""} {
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
if authorization != "" {
req.Header.Set("Authorization", authorization)
}
wantStatus(t, do(t, req), http.StatusNotFound)
wantLine(t, out.requestLines(t, i+1)[i], http.StatusNotFound,
requestlog.ActionAdmin)
}
if calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load())
}
}
func TestMetricsNeedTheTokenAndOtherPathsAreNotFound(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
for i, tc := range []struct {
method, path, authorization string
status int
}{
{http.MethodGet, proxy.MetricsPath, "", http.StatusUnauthorized},
{
http.MethodGet, proxy.MetricsPath, "Bearer " + strings.ToUpper(token),
http.StatusUnauthorized,
},
{http.MethodGet, proxy.MetricsPath, "Basic " + token, http.StatusUnauthorized},
{http.MethodGet, proxy.MetricsPath, bearer, http.StatusOK},
{http.MethodGet, proxy.MetricsPath, "bearer " + token, http.StatusOK},
{http.MethodPost, proxy.MetricsPath, bearer, http.StatusNotFound},
{http.MethodGet, proxy.MetricsPath + "/", bearer, http.StatusNotFound},
{http.MethodGet, "/_smallwebwaf/bans", bearer, http.StatusNotFound},
{http.MethodPost, proxy.HealthPath, "", http.StatusNotFound},
} {
req := newRequest(t, tc.method, addr, tc.path, http.NoBody)
if tc.authorization != "" {
req.Header.Set("Authorization", tc.authorization)
}
got := do(t, req)
wantStatus(t, got, tc.status)
wantLine(t, out.requestLines(t, i+1)[i], tc.status, requestlog.ActionAdmin)
if tc.status == http.StatusUnauthorized &&
got.header.Get("WWW-Authenticate") != "Bearer" {
t.Errorf("%q was answered without WWW-Authenticate: Bearer",
tc.authorization)
}
if tc.status == http.StatusOK &&
!strings.Contains(string(got.body), "# TYPE smallwebwaf_requests_total counter") {
t.Errorf("the metrics are\n%s", got.body)
}
}
if calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load())
}
}
func TestMetricsAreAskedForThroughTheChecks(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
metricsToken: token,
rateLimitPerMinute: "1",
})
// Asking for the metrics counts toward the client's limit of one
// request a minute, so its next request breaks it, and bans it. A
// banned client is refused the metrics too.
s.scrape(client)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
s.requestWithHeader(client, proxy.MetricsPath, "Authorization: "+bearer,
http.StatusForbidden, requestlog.ActionBanned)
}
func TestMetricsCountTheTraffic(t *testing.T) {
t.Parallel()
arrived, release := make(chan struct{}), make(chan struct{})
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path == "/held" {
close(arrived)
<-release
}
_, _ = io.WriteString(w, "hello")
})
releaseApp := sync.OnceFunc(func() { close(release) })
t.Cleanup(releaseApp)
addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token})
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
wantStatus(t, got, http.StatusOK)
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
out.requestLines(t, 2)
forward := `{action="forward",status_class="2xx"}`
notFound := `{action="admin",status_class="4xx"}`
// The request for the metrics is itself under way.
metrics := scrape(t, addr)
wantMetric(t, metrics, "smallwebwaf_requests_total"+forward, 1)
wantMetric(t, metrics, "smallwebwaf_requests_total"+notFound, 1)
wantMetric(t, metrics, "smallwebwaf_request_bytes_total"+forward, 3)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
float64(len("Not Found\n")))
wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
metric(t, metrics, "go_goroutines")
metric(t, metrics, "process_start_time_seconds")
// A request the app holds is under way until it ends.
httpClient := newClient(t)
held := newRequest(t, http.MethodGet, addr, "/held", http.NoBody)
ended := make(chan error, 1)
go func() {
res, err := httpClient.Do(held)
if err == nil {
err = readAnswer(res).err
}
ended <- err
}()
<-arrived
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
releaseApp()
err := <-ended
if err != nil {
t.Fatalf("held request: %v", err)
}
out.requestLines(t, 5)
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
}
func TestMetricsCountLimitsAndBans(t *testing.T) {
t.Parallel()
const (
scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
denied = "192.0.2.50" // in SWWAF_DENY_NETS
)
s, clk, _ := startWithClock(t, "", map[string]string{
metricsToken: token,
rateLimitPerMinute: "1",
rateLimitExemptNets: scraper,
denyNets: denied,
banResponse: "close",
limitBanDuration: "1h",
maxBanDuration: "2h",
})
// SWWAF_BAN_RESPONSE=close sends no status at all.
s.get(denied, 0, requestlog.ActionDenied)
// A first broken limit bans for an hour.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, 0, requestlog.ActionRateLimited)
metrics := s.scrape(scraper)
wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
clk.advance(time.Hour)
wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
// A limit broken again right after would ban for three hours, longer
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, 0, requestlog.ActionRateLimited)
metrics = s.scrape(scraper)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
// denied, client, and the scraper as of its earlier requests.
wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
}
func TestMetricsCountSizeAndTimeLimits(t *testing.T) {
t.Parallel()
// The app never answers /hang, so the timeout runs out however slowly
// the test runs.
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/hang" {
<-r.Context().Done()
}
})
addr, out := startProxy(t, app.URL, map[string]string{
metricsToken: token,
requestMaxBytes: sizeLimitSetting,
upstreamResponseTimeout: "100ms",
})
body := bytes.NewReader(make([]byte, 2*sizeLimit))
wantStatus(t, do(t, newRequest(t, http.MethodPost, addr, "/", body)),
http.StatusRequestEntityTooLarge)
wantStatus(t, get(t, addr, "/hang"), http.StatusGatewayTimeout)
out.requestLines(t, 2)
metrics := scrape(t, addr)
hits := "smallwebwaf_size_and_time_limit_hits_total"
wantMetric(t, metrics, hits+`{limit="SWWAF_REQUEST_MAX_BYTES"}`, 1)
wantMetric(t, metrics, hits+`{limit="SWWAF_UPSTREAM_RESPONSE_TIMEOUT"}`, 1)
}
func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
t.Parallel()
const fromFR = "198.51.100.20"
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
_, _ = io.WriteString(w, "hello")
})
env := map[string]string{
trustedProxies: trustLocalhost,
metricsToken: token,
metricsTopN: "2",
deniedCountries: "kp",
}
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
// The answers are kept before the requests, so that none waits for
// GeoJS.
server.GeoJS.Load([]lookup.Answer{
keptAnswer(fromKP, "KP"), keptAnswer(fromDE, "DE"), keptAnswer(fromFR, "FR"),
})
lines := 0
send := func(from string, times, status int) {
t.Helper()
for range times {
req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc"))
req.Header.Set(forwardedFor, from)
wantStatus(t, do(t, req), status)
// Each is counted before the next is sent, so that the
// countries are ranked in the order sent.
lines++
out.requestLines(t, lines)
}
}
// With two countries of their own, the third is counted as other.
send(fromKP, 3, http.StatusForbidden)
send(fromDE, 2, http.StatusOK)
send(fromFR, 1, http.StatusOK)
metrics := scrape(t, addr)
lines++
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
float64(3*len("Forbidden\n")))
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
float64(len("hello")))
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
// Once FR is busier than DE, it takes DE's place: its series counts
// from then on, and DE's is gone.
send(fromFR, 3, http.StatusOK)
metrics = scrape(t, addr)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
}
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
t.Parallel()
geojs := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(geojs.Close)
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{
trustedProxies: trustLocalhost,
metricsToken: token,
deniedCountries: "kp",
})
// GeoJS fails, so the client counts as coming from an unknown country,
// which SWWAF_DENIED_COUNTRIES does not refuse.
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, fromDE)
wantStatus(t, do(t, req), http.StatusOK)
// The client stops waiting for GeoJS after a second, so GeoJS's
// failure can come after its request has ended.
deadline := time.Now().Add(waitLimit)
metrics := scrape(t, addr)
for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
time.Now().Before(deadline) {
time.Sleep(pollInterval)
metrics = scrape(t, addr)
}
wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
}
// keptAnswer returns GeoJS's answer that the client at addr is in
// country, given now.
func keptAnswer(addr, country string) lookup.Answer {
now := time.Now()
return lookup.Answer{
Client: netip.MustParsePrefix(addr + "/32"), Country: country,
Answered: now, Used: now,
}
}
// scrape asks smallwebwaf at addr for the metrics, with the token, and
// returns them.
func scrape(t *testing.T, addr string) string {
t.Helper()
req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody)
req.Header.Set("Authorization", bearer)
got := do(t, req)
if got.status != http.StatusOK {
t.Fatalf("the metrics were answered %d", got.status)
}
return string(got.body)
}
// scrape asks for the metrics, with the token, from the client at from,
// and returns them.
func (s *sender) scrape(from string) string {
s.t.Helper()
_, metrics := s.requestWithHeader(from, proxy.MetricsPath, "Authorization: "+bearer,
http.StatusOK, requestlog.ActionAdmin)
return metrics
}
// metric returns the value of series in metrics, which are in the
// Prometheus text format. series is a name and its labels in the order of
// their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
// the test if there is no such series.
func metric(t *testing.T, metrics, series string) float64 {
t.Helper()
for line := range strings.Lines(metrics) {
value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ")
if !found {
continue
}
number, err := strconv.ParseFloat(value, 64)
if err != nil {
t.Fatalf("%s has the value %q", series, value)
}
return number
}
t.Fatalf("no series %s in the metrics:\n%s", series, metrics)
return 0
}
// wantMetric checks the value of series in metrics, as metric reads it.
func wantMetric(t *testing.T, metrics, series string, want float64) {
t.Helper()
got := metric(t, metrics, series)
if got != want {
t.Errorf("%s is %v, want %v", series, got, want)
}
}
// wantNoSeries checks that metrics have no series series.
func wantNoSeries(t *testing.T, metrics, series string) {
t.Helper()
if strings.Contains(metrics, "\n"+series+" ") {
t.Errorf("there is a series %s", series)
}
}
+28 -2
View File
@@ -8,11 +8,13 @@ import (
"log"
"log/slog"
"net/http"
"strings"
"time"
"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"
)
@@ -23,10 +25,18 @@ const (
appIdleConnTimeout = 90 * time.Second
)
// adminPrefix starts the path of every request for smallwebwaf itself,
// which never reaches the app.
const adminPrefix = "/_smallwebwaf/"
// HealthPath is smallwebwaf's health endpoint, which the container's
// health check asks.
const HealthPath = "/_smallwebwaf/healthz"
// MetricsPath is where the metrics are, for a request that carries
// SWWAF_METRICS_TOKEN.
const MetricsPath = "/_smallwebwaf/metrics"
// Params are what New needs.
type Params struct {
Config *config.Config
@@ -44,13 +54,14 @@ type Params struct {
}
// Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep.
// whose state the state files keep, and the metrics.
type Server struct {
*http.Server
Ledger *bans.Ledger
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
Metrics *metrics.Metrics
}
// New returns the server smallwebwaf runs: each request it reads passes
@@ -61,6 +72,7 @@ type Server struct {
// applies the timeouts and size limits from then on.
func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
m := metrics.New(params.Config.MetricsTopN)
h := &handler{
config: params.Config,
requestLog: params.RequestLog,
@@ -68,6 +80,7 @@ func New(params Params) *Server {
errorLog: errorLog,
transport: newTransport(),
now: params.Now,
metrics: m,
limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour,
@@ -83,8 +96,10 @@ func New(params Params) *Server {
URL: params.GeoJSURL,
Now: params.Now,
ProcessLog: params.ProcessLog,
Metrics: m,
}),
}
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
return &Server{
Server: &http.Server{
@@ -102,6 +117,7 @@ func New(params Params) *Server {
Ledger: h.ledger,
Limiter: h.limiter,
GeoJS: h.geojs,
Metrics: m,
}
}
@@ -114,6 +130,7 @@ type handler struct {
errorLog *log.Logger
transport http.RoundTripper
now func() time.Time
metrics *metrics.Metrics
limiter *ratelimit.Limiter
ledger *bans.Ledger
geojs *lookup.GeoJS
@@ -133,7 +150,8 @@ func newTransport() *http.Transport {
// ServeHTTP handles one request: it works out the client, runs the
// checks, passes the request to the app and the answer back within the
// limits, and writes the request's log line.
// limits, or answers it itself if it is for smallwebwaf, and writes the
// request's log line.
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
rq := h.newRequest(w, r)
defer rq.finish()
@@ -157,5 +175,13 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
// A request for smallwebwaf itself is answered where another would be
// passed to the app, so that it goes through every check first.
if strings.HasPrefix(r.URL.Path, adminPrefix) {
rq.answerAdmin()
return
}
rq.forward(r.Context())
}
+51 -17
View File
@@ -22,11 +22,13 @@ const flushAfterEachWrite time.Duration = -1
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
// the status the client is answered if the response has not started yet,
// 0 to close the connection without an answer, and the action the log
// line names.
// 0 to close the connection without an answer, the action the log line
// names, and the setting whose size or time limit the request passed, if
// that is why.
type refusal struct {
status int
action string
limit string
}
// request is one request on its way through smallwebwaf, from the moment
@@ -65,9 +67,11 @@ type request struct {
requestSent time.Time
}
// newRequest starts handling r: it notes the time and works out the
// client.
// newRequest starts handling r: it notes the time, counts the request as
// under way, and works out the client.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
h.metrics.RequestStarted()
start := time.Now()
peer := peerAddress(r)
trusted := h.config.TrustedProxies
@@ -141,6 +145,7 @@ func (rq *request) check(ctx context.Context) *refusal {
return &refusal{
status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge,
limit: "SWWAF_REQUEST_MAX_BYTES",
}
}
@@ -206,7 +211,11 @@ func (rq *request) modifyResponse(res *http.Response) error {
maxBytes := rq.h.config.ResponseMaxBytes
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge})
rq.refuse(refusal{
status: http.StatusBadGateway,
action: requestlog.ActionTooLarge,
limit: "SWWAF_RESPONSE_MAX_BYTES",
})
return errResponseTooLarge
}
@@ -278,7 +287,8 @@ func (rq *request) refuse(r refusal) {
rq.cancel()
}
// finish ends the request's timeouts and writes its log line.
// finish ends the request's timeouts, counts it in the metrics and writes
// its log line.
func (rq *request) finish() {
rq.stopTimers()
@@ -295,24 +305,37 @@ func (rq *request) finish() {
line.RequestBytes = rq.body.bytes.Load()
}
// limit is the setting whose size or time limit the request passed.
var limit string
switch {
case refused != nil:
line.Action = refused.action
limit = refused.limit
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
// take the response.
line.Action = requestlog.ActionTimedOut
limit = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
line.Aborted = true
}
now := time.Now()
line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start))
duration := now.Sub(rq.start)
line.DurationTotal = requestlog.Milliseconds(duration)
var upstreamDuration time.Duration
if !rq.upstreamStart.IsZero() {
line.DurationUpstreamTotal = requestlog.Milliseconds(now.Sub(rq.upstreamStart))
upstreamDuration = now.Sub(rq.upstreamStart)
line.DurationUpstreamTotal = requestlog.Milliseconds(upstreamDuration)
}
// Counted before the log line is written, so that the metrics count
// every request whose line is out.
rq.h.metrics.RequestEnded(line, limit, duration, upstreamDuration)
err := requestlog.Write(rq.h.requestLog, line)
if err != nil {
rq.h.processLog.Error("writing the request log failed", "error", err.Error())
@@ -327,9 +350,12 @@ func (rq *request) addToHistory() {
requestBytes = rq.body.bytes.Load()
}
forwarded := !rq.upstreamStart.IsZero()
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country,
Forwarded: !rq.upstreamStart.IsZero(),
Forwarded: forwarded,
Refused: !forwarded && rq.refused.Load() != nil,
Status: rq.out.status,
RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes,
@@ -369,21 +395,26 @@ func (rq *request) startRequestTimers() {
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
rq.clientRequestTimer = time.AfterFunc(
time.Until(rq.clientRequestDeadline()), rq.requestTimedOut)
time.Until(rq.clientRequestDeadline()), func() {
rq.requestTimedOut("SWWAF_CLIENT_REQUEST_TIMEOUT")
})
}
timeout := rq.h.config.UpstreamRequestTimeout
if timeout > 0 {
rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut)
rq.upstreamRequestTimer = time.AfterFunc(timeout, func() {
rq.requestTimedOut("SWWAF_UPSTREAM_REQUEST_TIMEOUT")
})
}
}
// requestTimedOut is called when a request timeout runs out while the
// request is still on its way to the app. The answer names the side
// smallwebwaf was waiting on at that moment: 408 when it was waiting for
// the client to send more of its body, 504 when it was waiting for the
// app to be reached or to take what it had.
func (rq *request) requestTimedOut() {
// requestTimedOut is called when limit, SWWAF_CLIENT_REQUEST_TIMEOUT or
// SWWAF_UPSTREAM_REQUEST_TIMEOUT, runs out while the request is still on
// its way to the app. The answer names the side smallwebwaf was waiting
// on at that moment: 408 when it was waiting for the client to send more
// of its body, 504 when it was waiting for the app to be reached or to
// take what it had.
func (rq *request) requestTimedOut(limit string) {
rq.mu.Lock()
defer rq.mu.Unlock()
@@ -395,6 +426,7 @@ func (rq *request) requestTimedOut() {
rq.refuse(refusal{
status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut,
limit: limit,
})
return
@@ -403,6 +435,7 @@ func (rq *request) requestTimedOut() {
rq.refuse(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
limit: limit,
})
// The transport gives up on the app only once its Read of the
// client's body returns, so that Read is ended now. The lock keeps
@@ -452,6 +485,7 @@ func (rq *request) responseTimedOut() {
rq.refuse(refusal{
status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut,
limit: "SWWAF_UPSTREAM_RESPONSE_TIMEOUT",
})
}
}
+8 -5
View File
@@ -19,26 +19,29 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5},
{Country: "FR", Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response.
{Status: 0},
{Refused: true, Status: 0},
// Answered at smallwebwaf's own endpoints: neither forwarded nor
// refused.
{Status: 404},
} {
limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
}
want := ratelimit.History{
FirstSeen: start,
LastSeen: start.Add(5 * time.Minute),
LastSeen: start.Add(6 * time.Minute),
Country: "FR",
LookedUp: start.Add(3 * time.Minute),
Requests: 6,
Requests: 7,
Forwarded: 4,
Refused: 2,
RequestBytes: 15,
ResponseBytes: 122,
Responses: ratelimit.Responses{
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 1, Status5xx: 1,
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 2, Status5xx: 1,
},
Offences: ratelimit.Offences{Limit: 1},
}
+17 -4
View File
@@ -71,7 +71,8 @@ type History struct {
Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to
// the app, Refused those refused before anything reached it.
// the app, Refused those refused before anything reached it, and
// neither those smallwebwaf answered at its own endpoints.
Requests int64 `json:"requests"`
Forwarded int64 `json:"forwarded"`
Refused int64 `json:"refused"`
@@ -103,9 +104,11 @@ type Offences struct {
type Request struct {
// Country is the client's country, when the request looked it up.
Country string
// Forwarded is true for a request passed to the app, false for one
// refused before anything reached it.
// Forwarded is true for a request passed to the app, Refused for one
// refused before anything reached it. Both are false for a request
// smallwebwaf answered at its own endpoints.
Forwarded bool
Refused bool
// Status is what the client was sent, 0 if nothing was.
Status int
// RequestBytes and ResponseBytes are the body bytes of the request
@@ -199,7 +202,9 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
h.Requests++
if r.Forwarded {
h.Forwarded++
} else {
}
if r.Refused {
h.Refused++
}
@@ -235,6 +240,14 @@ func (l *Limiter) Requests(netblock netip.Prefix) int64 {
return requests
}
// Len returns how many clients are in the table.
func (l *Limiter) Len() int {
l.mu.Lock()
defer l.mu.Unlock()
return l.clients.Len()
}
// Snapshot returns every client in the table, sorted by address, as
// clients.json lists them.
func (l *Limiter) Snapshot() []Client {
+1
View File
@@ -89,6 +89,7 @@ func Run(ctx context.Context, params Params) int {
GeoJS: server.GeoJS,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
if err != nil {
processLog.Error("cannot use the state files", "error", err.Error())
+22
View File
@@ -121,6 +121,28 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
}
}
func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) {
t.Parallel()
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
out := &output{}
status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
}
line := out.line(t, "msg", "invalid setting")
if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" {
t.Errorf("start refused with %v", line)
}
if strings.Contains(out.text(), token) {
t.Errorf("the output shows the token:\n%s", out.text())
}
}
func TestAddressInUseStopsTheStart(t *testing.T) {
t.Parallel()
+15 -3
View File
@@ -21,6 +21,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
@@ -62,6 +63,8 @@ type Params struct {
Now func() time.Time
// ProcessLog receives what was read, and the writes that fail.
ProcessLog *slog.Logger
// Metrics count each file's writes.
Metrics *metrics.Metrics
}
// Files are the state files of a running smallwebwaf.
@@ -204,7 +207,7 @@ func (f *Files) writeBans() error {
return fmt.Errorf("encode %s: %w", bansJSON, err)
}
return write(f.params.Dir, bansJSON, append(data, '\n'))
return f.writeCounted(bansJSON, append(data, '\n'))
}
// writeClients writes clients.json.
@@ -214,7 +217,7 @@ func (f *Files) writeClients() error {
return fmt.Errorf("encode %s: %w", clientsJSON, err)
}
return write(f.params.Dir, clientsJSON, data)
return f.writeCounted(clientsJSON, data)
}
// writeLookups writes lookups.json.
@@ -224,7 +227,16 @@ func (f *Files) writeLookups() error {
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
}
return write(f.params.Dir, lookupsJSON, data)
return f.writeCounted(lookupsJSON, data)
}
// writeCounted writes data to the state file name, as write does, and
// counts the write in the metrics.
func (f *Files) writeCounted(name string, data []byte) error {
err := write(f.params.Dir, name, data)
f.params.Metrics.StateFileWritten(name, len(data), err)
return err
}
// newBanEntry returns ban as bans.json holds it.
+114 -2
View File
@@ -4,10 +4,13 @@ import (
"context"
"encoding/json"
"log/slog"
"net/http"
"net/http/httptest"
"net/netip"
"os"
"path/filepath"
"slices"
"strconv"
"strings"
"testing"
"testing/synctest"
@@ -15,6 +18,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/state"
)
@@ -402,6 +406,59 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
}
}
func TestWritesAreCountedInTheMetrics(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
const (
ofBans = `{file="bans.json"}`
ofClients = `{file="clients.json"}`
)
got := scrape(t, params)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 1)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofClients, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 0)
wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans,
float64(len(permanentBansJSON)))
written := metric(t, got, "smallwebwaf_state_file_last_write_timestamp_seconds"+ofBans)
if written < float64(time.Now().Add(-time.Hour).Unix()) {
t.Errorf("bans.json was last written at %v, not by that write", written)
}
// A directory in the way of bans.json's temporary file fails its next
// write, which leaves its size as it was, although it has a ban more.
err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
err = files.WriteAll()
if err == nil {
t.Fatal("the write did not fail")
}
got = scrape(t, params)
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 2)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofClients, 0)
wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans,
float64(len(permanentBansJSON)))
}
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
t.Parallel()
@@ -431,6 +488,7 @@ func midnight() time.Time {
// hold nothing yet. GeoJS is never asked.
func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler)
m := metrics.New(1)
return state.Params{
Dir: dir,
@@ -442,10 +500,13 @@ func newParams(dir string) state.Params {
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
}),
Limiter: ratelimit.New(ratelimit.Limits{}),
GeoJS: lookup.New(lookup.Params{Now: midnight, ProcessLog: discard}),
Limiter: ratelimit.New(ratelimit.Limits{}),
GeoJS: lookup.New(lookup.Params{
Now: midnight, ProcessLog: discard, Metrics: m,
}),
Now: midnight,
ProcessLog: discard,
Metrics: m,
}
}
@@ -633,3 +694,54 @@ func wantEntries(t *testing.T, path, key string, want ...string) {
}
}
}
// scrape returns the metrics of params, in the Prometheus text format.
func scrape(t *testing.T, params state.Params) string {
t.Helper()
recorder := httptest.NewRecorder()
params.Metrics.ServeHTTP(recorder,
httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody))
if recorder.Code != http.StatusOK {
t.Fatalf("the metrics were answered %d", recorder.Code)
}
return recorder.Body.String()
}
// metric returns the value of series in text, the metrics, such as
// smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test
// if there is no such series.
func metric(t *testing.T, text, series string) float64 {
t.Helper()
for line := range strings.Lines(text) {
value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ")
if !found {
continue
}
number, err := strconv.ParseFloat(value, 64)
if err != nil {
t.Fatalf("%s has the value %q", series, value)
}
return number
}
t.Fatalf("no series %s in the metrics:\n%s", series, text)
return 0
}
// wantMetric checks the value of series in text, the metrics, as metric
// reads it.
func wantMetric(t *testing.T, text, series string, want float64) {
t.Helper()
got := metric(t, text, series)
if got != want {
t.Errorf("%s is %v, want %v", series, got, want)
}
}