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/ratelimit" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) const errorBurstThreshold = "SWWAF_ERROR_BURST_THRESHOLD" // refused is a request the tests here send, which smallwebwaf refuses // after a rule file match, or for a missing or wrong token. type refused int const ( // blockRule is a request testRules' block rule refuses with 403. blockRule refused = iota // banRule is one its ban rule refuses with 403, and bans the client // for. banRule // noMetricsToken is one for the metrics without a token, and // wrongAdminToken one for the bans with the metrics token, each // refused with 401. noMetricsToken wrongAdminToken ) // send sends r from the client at from, checks its answer and log line as // sender.request does, and returns the line. func (r refused) send(s *sender, from string) logLine { s.t.Helper() switch r { case blockRule: return s.request(from, blockedPath, http.StatusForbidden, requestlog.ActionRuleBlocked) case banRule: return s.request(from, probePath, http.StatusForbidden, requestlog.ActionBanned) case noMetricsToken: return s.request(from, proxy.MetricsPath, http.StatusUnauthorized, requestlog.ActionAdmin) case wrongAdminToken: line, _ := s.requestWithHeader(from, proxy.BansPath, "Authorization: "+bearer, http.StatusUnauthorized, requestlog.ActionAdmin) return line } s.t.Fatalf("no request for the refusal %d", r) return logLine{} } // startForErrorBurst is startWithClock with testRules, both tokens and // SWWAF_ERROR_BURST_THRESHOLD at threshold, and the settings in env. func startForErrorBurst( t *testing.T, threshold string, env map[string]string, ) (*sender, *clock, *proxy.Server) { t.Helper() settings := map[string]string{ errorBurstThreshold: threshold, rulesDir: writeRules(t, testRules), adminToken: adminSecret, metricsToken: token, } maps.Copy(settings, env) return startWithClock(t, "", settings) } func TestErrorBurstBreaksAtOneOverTheThreshold(t *testing.T) { t.Parallel() for _, tc := range []struct { name string // refusals are four, one over the threshold of three. refusals []refused }{ {"block rule", []refused{blockRule, blockRule, blockRule, blockRule}}, { "missing or wrong token", []refused{noMetricsToken, wrongAdminToken, noMetricsToken, wrongAdminToken}, }, { "a mix ending in a ban rule", []refused{blockRule, noMetricsToken, blockRule, banRule}, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() s, _, _ := startForErrorBurst(t, "3", nil) // Three refusals break nothing, and the app's answers between // them are not counted. for i, r := range tc.refusals[:3] { line := r.send(s, client) if line.LimitHit != "" || line.Offence != "" { t.Errorf("refusal %d: log line has limit_hit %q and offence %q, "+ "want none", i+1, line.LimitHit, line.Offence) } s.get(client, http.StatusOK, requestlog.ActionForward) } // The fourth is answered as the others were, breaks the error // burst, and bans the client. line := tc.refusals[3].send(s, client) if line.LimitHit != requestlog.LimitHitErrorBurst || line.Offence != requestlog.OffenceLimit { t.Errorf("log line has limit_hit %q and offence %q, want error_burst "+ "and limit", line.LimitHit, line.Offence) } s.get(client, http.StatusForbidden, requestlog.ActionBanned) }) } } func TestErrorBurstBanNotesHistoryAndMetrics(t *testing.T) { t.Parallel() const scraper = "192.0.2.200" s, clk, server := startForErrorBurst(t, "2", nil) start := clk.Now() blockRule.send(s, client) wrongAdminToken.send(s, client) line := blockRule.send(s, client) expires := start.Add(time.Hour) if line.BanExpires != requestlog.FormatTime(expires) { t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires) } netblock := netip.MustParsePrefix(client + "/32") want := bans.Ban{ Netblock: netblock, Start: start, Expires: expires, Cause: bans.CauseLimit, Reason: "refusals per minute over the limit of 2", Notes: bans.Notes{ Kind: ratelimit.KindRefusals, Limit: 2, Window: minute, Count: 3, Request: bans.Request{ Time: start, Method: http.MethodGet, Host: appHost, Path: blockedPath, Status: http.StatusForbidden, UserAgent: userAgent, }, Requests: 3, }, } got := server.Ledger.Bans(netblock) if len(got) != 1 || !reflect.DeepEqual(got[0], want) { t.Fatalf("bans\n%+v\nwant\n%+v", got, want) } wantOffences := ratelimit.Offences{Limit: 1, RuleBlocked: 2, TokenRefused: 1} if offences := historyOf(t, server, client).Offences; offences != wantOffences { t.Errorf("history counts the offences %+v, want %+v", offences, wantOffences) } metrics := s.scrape(scraper) wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+ `kind="refusals",window="minute"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="rule_blocked"}`, 2) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="token_refused"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1) } func TestErrorBurstIsNotLoweredForAClientWithLowerLimits(t *testing.T) { t.Parallel() geojsURL, _ := startGeoJS(t) s, _, server := startWithClock(t, geojsURL, map[string]string{ errorBurstThreshold: "2", rulesDir: writeRules(t, testRules), countryLimitPercent: countryDEHalf, }) // Half of the threshold would be one, which the second refusal is over. for range 2 { line := blockRule.send(s, fromDE) if line.LimitHit != "" { t.Errorf("log line has limit_hit %q, want none", line.LimitHit) } } blockRule.send(s, fromDE) got := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32")) if len(got) != 1 || got[0].Notes.Limit != 2 || got[0].Notes.LimitPercent != nil { t.Errorf("bans %+v, want one for the limit of 2, without a limit percentage", got) } } func TestErrorBurstDoesNotCountTheAppsAnswers(t *testing.T) { t.Parallel() statuses := map[string]int{ "/missing": http.StatusNotFound, "/private": http.StatusUnauthorized, "/forbidden": http.StatusForbidden, } s, _, _, queue := startAppWithAlerts(t, func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(statuses[r.URL.Path]) }, map[string]string{errorBurstThreshold: "1", rulesDir: writeRules(t, testRules)}) for range 2 { for path, status := range statuses { s.request(client, path, status, requestlog.ActionForward) } } // The first refusal is one, not over the threshold. line := blockRule.send(s, client) if line.LimitHit != "" { t.Errorf("log line has limit_hit %q, want none", line.LimitHit) } // No ban was made, nor its alert raised. s.request(client, "/missing", http.StatusNotFound, requestlog.ActionForward) wantAlerts(t, queue) } func TestErrorBurstOffOrAtItsDefault(t *testing.T) { t.Parallel() const off = "off" for _, tc := range []struct { threshold string // broken is whether the 31st refusal breaks the error burst. broken bool }{ {"", true}, {off, false}, } { t.Run(errorBurstThreshold+"="+tc.threshold, func(t *testing.T) { t.Parallel() env := map[string]string{rulesDir: writeRules(t, testRules)} if tc.threshold != "" { env[errorBurstThreshold] = tc.threshold } s, _, _ := startWithClock(t, "", env) var line logLine for range 31 { line = blockRule.send(s, client) } if broken := line.LimitHit == requestlog.LimitHitErrorBurst; broken != tc.broken { t.Errorf("the 31st refusal broke the error burst: %t, want %t", broken, tc.broken) } }) } } func TestErrorBurstCountsEachClientTheChecksApplyTo(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, _, _ := startForErrorBurst(t, "1", map[string]string{ allowNets: allowed, rateLimitExemptNets: exempt, }) // A client in SWWAF_ALLOW_NETS still needs the token, but is not // counted. for range 3 { line := noMetricsToken.send(s, allowed) if line.LimitHit != "" { t.Errorf("log line has limit_hit %q, want none", line.LimitHit) } } // One the rate limits do not apply to is. noMetricsToken.send(s, exempt) line := wrongAdminToken.send(s, exempt) if line.LimitHit != requestlog.LimitHitErrorBurst { t.Errorf("log line has limit_hit %q, want error_burst", line.LimitHit) } s.get(exempt, http.StatusForbidden, requestlog.ActionBanned) } func TestErrorBurstBanSetsTheRefusalsBackToZero(t *testing.T) { t.Parallel() s, clk, _ := startForErrorBurst(t, "1", map[string]string{limitBanDuration: "1s"}) blockRule.send(s, client) blockRule.send(s, client) // Within the same minute, once the ban has ended, the next refusal is // the first again. clk.advance(time.Second) line := blockRule.send(s, client) if line.LimitHit != "" { t.Errorf("log line has limit_hit %q, want none", line.LimitHit) } } func TestObserveModeLogsAndAlertsTheErrorBurst(t *testing.T) { t.Parallel() s, clk, server, queue := startWithAlerts(t, map[string]string{ mode: observe, errorBurstThreshold: "1", rulesDir: writeRules(t, testRules), adminToken: adminSecret, }) start := clk.Now() held := bans.Ban{ Netblock: netip.MustParsePrefix(otherClient + "/32"), Start: start, Expires: start.Add(time.Hour), Cause: bans.CauseAdmin, } server.Ledger.Load([]bans.Ban{held}) // Under a ban, enforce mode would have refused these before the // endpoint, so their tokens are not counted. for range 2 { line := wrongAdminToken.send(s, otherClient) wantWouldAction(t, line, requestlog.ActionBanned) if line.LimitHit != "" { t.Errorf("log line has limit_hit %q, want none", line.LimitHit) } } // The block rule's refusal, which enforce mode would have answered 403, // is the second of the client's, and would have banned it. wrongAdminToken.send(s, client) line := s.request(client, blockedPath, http.StatusOK, requestlog.ActionForward) wantWouldAction(t, line, requestlog.ActionRuleBlocked) if line.LimitHit != requestlog.LimitHitErrorBurst || line.BanExpires != "" { t.Errorf("log line has limit_hit %q and ban_expires %q, want error_burst "+ "and none", line.LimitHit, line.BanExpires) } if got := server.Ledger.Snapshot(); len(got) != 1 || !reflect.DeepEqual(got[0], held) { t.Errorf("bans %+v, want only the one held", got) } waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] if len(waiting) != 1 { t.Fatalf("%d alerts wait, want 1: %+v", len(waiting), waiting) } notes, _ := waiting[0].Detail["notes"].(bans.Notes) if notes.Kind != ratelimit.KindRefusals || notes.Count != 2 || notes.Request.Status != http.StatusForbidden { t.Errorf("the alert's notes are %+v, want two refusals, the last answered 403", notes) } alert := banAlert(alerts.EventBan, start, client, bans.Ban{ Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit, Reason: "refusals per minute over the limit of 1", Notes: notes, }, requestlog.FormatTime(start.Add(time.Hour))) alert.Detail["mode"] = observe wantAlerts(t, queue, alert) }