package proxy_test import ( "maps" "net/http" "net/netip" "reflect" "slices" "strconv" "sync/atomic" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/anomaly" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // The anomaly thresholds: the prefix of a scope followed by the end of a // count. const ( anomalyClient = "SWWAF_ANOMALY_CLIENT_" anomalyNet = "SWWAF_ANOMALY_NET_" anomalyASN = "SWWAF_ANOMALY_ASN_" anomalyTotal = "SWWAF_ANOMALY_TOTAL_" anomalyWatch = "SWWAF_WATCH_" requestsPerMinute = "REQUESTS_PER_MINUTE" requestsPerHour = "REQUESTS_PER_HOUR" bytesPerMinute = "BYTES_PER_MINUTE" bytesPerHour = "BYTES_PER_HOUR" ) // The other anomaly settings. const ( anomalyNetV4Prefix = "SWWAF_ANOMALY_NET_V4_PREFIX" anomalyNetV6Prefix = "SWWAF_ANOMALY_NET_V6_PREFIX" watchNets = "SWWAF_WATCH_NETS" ) const ( // clientsNet is the netblock around client at the default length, and // office a named netblock of the same. clientsNet = "203.0.113.0/24" office = "office=" + clientsNet // aLot is a threshold no test reaches. aLot = "1000" // hour is the window an alert names for a threshold per hour. hour = "hour" ) func TestEachScopeAndWindowOverItsThresholdAlertsOncePerCooldown(t *testing.T) { t.Parallel() for _, scope := range []struct { prefix, scope string // netblock is the alert's, and counted what its reason names. extra // is what its detail gives besides what every anomaly alert's does. netblock netip.Prefix counted string extra map[string]any }{ { anomalyClient, anomaly.ScopeClient, netip.MustParsePrefix(client + "/32"), "the client " + client + "/32", nil, }, { anomalyNet, anomaly.ScopeNet, netip.MustParsePrefix(clientsNet), "the netblock " + clientsNet, nil, }, {anomalyASN, anomaly.ScopeASN, netip.Prefix{}, asnDE, map[string]any{"asn": asnDE}}, {anomalyTotal, anomaly.ScopeTotal, netip.Prefix{}, "the whole service", nil}, { anomalyWatch, anomaly.ScopeWatch, netip.MustParsePrefix(clientsNet), "the named netblock office, " + clientsNet, map[string]any{"name": "office"}, }, } { for _, threshold := range []struct { end, kind, window string // value is the threshold, which the third upload of 100 bytes // takes the count over, to count. value int64 count float64 }{ {requestsPerMinute, ratelimit.KindRequests, minute, 2, 3}, {requestsPerHour, ratelimit.KindRequests, hour, 2, 3}, {bytesPerMinute, ratelimit.KindBytes, minute, 250, 300}, {bytesPerHour, ratelimit.KindBytes, hour, 250, 300}, } { setting := scope.prefix + threshold.end value := strconv.FormatInt(threshold.value, 10) t.Run(setting, func(t *testing.T) { t.Parallel() s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{ setting: value, watchNets: office, }) start := clk.Now() // The third upload takes the count over the threshold, and the // fourth, within the cooldown, is held back. Each is passed to // the app. for range 4 { s.uploadFrom(client) } detail := map[string]any{ "scope": scope.scope, "window": threshold.window, "kind": threshold.kind, "count": threshold.count, "threshold": threshold.value, } maps.Copy(detail, scope.extra) wantAlerts(t, queue, alerts.Alert{ Instance: alertInstance, Time: start, Event: alerts.EventAnomaly, Client: netip.MustParseAddr(client), Netblock: scope.netblock, ASN: asnDE, ASName: asNameDE, Country: "DE", Reason: threshold.kind + " per " + threshold.window + " of " + scope.counted + " over the threshold of " + value, Detail: detail, }) wantAlertedAgainOnceTheCooldownHasRunOut(t, s, clk, queue) if held := server.Ledger.Snapshot(); len(held) != 0 { t.Errorf("the ledger holds %+v, want no ban", held) } }) } } } // wantAlertedAgainOnceTheCooldownHasRunOut checks that, once the cooldown // has run out after a first alert, which held back one repeat, the next // count over the threshold, at the latest three uploads from client on, // raises another alert, giving that repeat. func wantAlertedAgainOnceTheCooldownHasRunOut( t *testing.T, s *sender, clk *clock, queue *alerts.Queue, ) { t.Helper() clk.advance(15 * time.Minute) for range 3 { s.uploadFrom(client) } waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] if len(waiting) != 2 || !waiting[1].Time.Equal(clk.Now()) || waiting[1].SuppressedRepeats != 1 { t.Errorf("alerts wait %+v, want the first and another, with 1 repeat", waiting) } } func TestEveryRequestIsCountedWhateverIsDoneWithIt(t *testing.T) { t.Parallel() const ( allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS denied = "192.0.2.20" // in SWWAF_DENY_NETS ) s, _, _, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{ anomalyClient + requestsPerMinute: "2", allowNets: allowed, rateLimitExemptNets: exempt, rateLimitExemptPaths: "/static/", denyNets: denied, }) // The third request of each takes its client's count over the threshold // of 2. for _, sent := range []struct { from, path string status int action string }{ {allowed, "/", http.StatusOK, requestlog.ActionForward}, {exempt, "/", http.StatusOK, requestlog.ActionForward}, {client, "/static/app.js", http.StatusOK, requestlog.ActionForward}, {denied, "/", http.StatusForbidden, requestlog.ActionDenied}, } { for range 3 { s.request(sent.from, sent.path, sent.status, sent.action) } } waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] got := make([]string, 0, len(waiting)) for _, alert := range waiting { got = append(got, alert.Client.String()) } if want := []string{allowed, exempt, client, denied}; !slices.Equal(got, want) { t.Errorf("alerts for the clients %v, want %v", got, want) } } func TestThresholdsOffCountNothingAndAlertNothing(t *testing.T) { t.Parallel() // With every threshold off, nothing is counted. s, server, queue := startWithLookups(t, map[string]string{watchNets: office}) for range 5 { s.uploadFrom(client) } if counters := server.Anomalies.Snapshot(); len(counters) != 0 { t.Errorf("counters %+v, want none", counters) } wantAlerts(t, queue) // With one set, its count alone is counted, in its scope alone. s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{ anomalyNet + requestsPerMinute: aLot, watchNets: office, }) for range 5 { s.uploadFrom(client) } want := []anomaly.Counter{{ Scope: anomaly.ScopeNet, Netblock: netip.MustParsePrefix(clientsNet), Minute: ratelimit.Buckets{Start: clk.Now(), Current: 5}, }} if got := server.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) { t.Errorf("counters\n%+v\nwant\n%+v", got, want) } wantAlerts(t, queue) } func TestNetblockAroundAClientIsAsLongAsTheSettingsSay(t *testing.T) { t.Parallel() for _, tc := range []struct { name string env map[string]string // Each client of sent sends one request, and want gives the // netblocks they are counted in, each with its requests. sent []string want map[string]int64 }{ { "by default", nil, []string{client, "203.0.113.200", "192.0.2.7", ipv6Client, "2001:db8:0:ffff::1"}, map[string]int64{clientsNet: 2, "192.0.2.0/24": 1, "2001:db8::/48": 2}, }, { "as set", map[string]string{anomalyNetV4Prefix: "16", anomalyNetV6Prefix: "32"}, []string{client, "203.0.200.1", ipv6Client, "2001:db8:ffff::1"}, map[string]int64{"203.0.0.0/16": 2, "2001:db8::/32": 2}, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() env := map[string]string{anomalyNet + requestsPerMinute: aLot} maps.Copy(env, tc.env) s, _, server, _ := startAppWithAlerts(t, readAndAnswer, env) for _, from := range tc.sent { s.get(from, http.StatusOK, requestlog.ActionForward) } got := map[string]int64{} for _, counter := range server.Anomalies.Snapshot() { got[counter.Netblock.String()] = counter.Minute.Current } if !maps.Equal(got, tc.want) { t.Errorf("requests by netblock %v, want %v", got, tc.want) } }) } } func TestClientIsCountedForItsASNumberOnceTheLookupGivesOne(t *testing.T) { t.Parallel() s, server, _ := startWithLookups(t, map[string]string{ anomalyASN + requestsPerMinute: aLot, }) // The lookup database does not hold unplaced. for _, from := range []string{fromDE, fromDE, fromKP, noCountry, unplaced} { s.uploadFrom(from) } got := map[string]int64{} for _, counter := range server.Anomalies.Snapshot() { got[counter.ASN] = counter.Minute.Current } if want := map[string]int64{asnDE: 2, asnKP: 1, "AS64500": 1}; !maps.Equal(got, want) { t.Errorf("requests by AS number %v, want %v", got, want) } } func TestRequestCountsForTheASNumberGeoJSGivesBeforeItEnds(t *testing.T) { t.Parallel() // The stand-in for GeoJS answers only once released, which the app // does as it answers the request, and then waits until the answer is // kept. geojsURL, _, release := startHeldGeoJS(t) var server atomic.Pointer[proxy.Server] app := startApp(t, func(http.ResponseWriter, *http.Request) { release() waitUntil(func() bool { _, kept := server.Load().GeoJS.Kept(netip.MustParsePrefix(fromDE + "/32")) return kept }) }) clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)} addr, out, started := startProxyWithClock(t, app.URL, geojsURL, clk.Now, map[string]string{ trustedProxies: trustLocalhost, lookupTimeout: "1h", anomalyASN + requestsPerMinute: aLot, }) server.Store(started) // The request went on without the answer, and is counted for the AS // number it gives. s := &sender{t: t, addr: addr, out: out} if line := s.get(fromDE, http.StatusOK, requestlog.ActionForward); line.ASN != "" { t.Errorf("log line has AS number %q, want none: the request waited", line.ASN) } want := []anomaly.Counter{{ Scope: anomaly.ScopeASN, ASN: asnDE, Minute: ratelimit.Buckets{Start: clk.Now(), Current: 1}, }} if got := started.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) { t.Errorf("counters\n%+v\nwant\n%+v", got, want) } } func TestEachNamedNetblockCountsTheClientsInIt(t *testing.T) { t.Parallel() s, _, server, _ := startAppWithAlerts(t, readAndAnswer, map[string]string{ anomalyWatch + requestsPerMinute: aLot, watchNets: office + ",wide=203.0.0.0/16,other=198.51.100.0/25", }) // client is in office and in wide. for _, from := range []string{client, "203.0.200.1", "192.0.2.7"} { s.get(from, http.StatusOK, requestlog.ActionForward) } got := map[string]int64{} for _, counter := range server.Anomalies.Snapshot() { got[counter.Name] = counter.Minute.Current } if want := map[string]int64{"office": 1, "wide": 2}; !maps.Equal(got, want) { t.Errorf("requests by named netblock %v, want %v", got, want) } }