package reputation_test import ( "bytes" "context" "encoding/binary" "errors" "io" "log/slog" "net" "net/http" "net/http/httptest" "net/netip" "reflect" "slices" "strings" "sync" "testing" "testing/synctest" "time" "sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/reputation" ) // The tests of the DNSBL zones run in synctest bubbles, as those of the // lists do, and the resolver the zones are asked through is a stand-in // reached through an in-memory connection, net.Pipe's, for the same // reason. They run one at a time, none in parallel with another test of // this package: Go's resolver counts the queries under way in one // sync.WaitGroup for the whole process, and the process fails when // queries from two bubbles, or from a bubble and from outside one, are // under way at once. TestMain has the resolver make its configuration, // which it makes on its first query, outside every bubble, since the // configuration holds a channel, which the bubble it was made in would // keep to itself. const ( // zone and otherZone are the DNSBL zones the tests name. zone = "dnsbl.example" otherZone = "other.example" // cacheTTL is the tests' SWWAF_REPUTATION_CACHE_TTL, and timeout their // SWWAF_REPUTATION_TIMEOUT: a second, the least time /etc/resolv.conf // can have Go's resolver wait for one server, so that it is the // DNSBL's own timeout that ends a query, whatever that file says. cacheTTL = 24 * time.Hour timeout = time.Second // listed and unlisted are clients zone is asked about by the names // listedName and unlistedName, and most tests have zone list the first // alone, by answering with listing. listed = "192.0.2.99" unlisted = "192.0.2.100" listedName = "99.2.0.192." + zone + "." unlistedName = "100.2.0.192." + zone + "." listing = "127.0.0.2" ) // The DNS response codes the stand-in answers with, besides no error. const ( serverFailure = 2 noSuchName = 3 refused = 5 ) var errNoNetwork = errors.New("the test dials nothing") func TestMain(m *testing.M) { // A query that fails at once, as nothing is dialled for it. resolver := &net.Resolver{ PreferGo: true, Dial: func(context.Context, string, string) (net.Conn, error) { return nil, errNoNetwork }, } _, _ = resolver.LookupNetIP(context.Background(), "ip4", "warm-up.invalid.") m.Run() } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestZonesListOrNotClientsByTheirIPv4AndIPv6Addresses(t *testing.T) { synctest.Test(t, func(t *testing.T) { // The addresses of the examples of RFC 5782, and the names it // gives for them. const ( v4 = "192.0.2.99" v6 = "2001:db8:1:2:3:4:567:89ab" // v6Name is the hex digits of v6, in reverse order. v6Name = "b.a.9.8.7.6.5.0.4.0.0.0.3.0.0.0.2.0.0.0.1.0.0.0.8.b.d.0.1.0.0.2." ) resolver := &resolverStandIn{answers: map[string]answer{ "99.2.0.192." + zone + ".": {addrs: []string{listing}}, v6Name + otherZone + ".": {addrs: []string{"127.0.0.4", "127.0.0.10"}}, "99.2.0.192." + otherZone + ".": {}, }} dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone)) // Neither client has a verdict yet, so neither is listed, and each // zone is asked about each. wantZones(t, dnsbl, v4) wantZones(t, dnsbl, v6) synctest.Wait() wantZones(t, dnsbl, v4, zone) wantZones(t, dnsbl, v6, otherZone) wantAsked(t, resolver, "99.2.0.192."+zone+".", "99.2.0.192."+otherZone+".", v6Name+zone+".", v6Name+otherZone+".") }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestListedByNeverWaitsForAQuery(t *testing.T) { synctest.Test(t, func(t *testing.T) { dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone)) began := time.Now() // The second, while the first's query is under way, starts none. wantZones(t, dnsbl, listed) wantZones(t, dnsbl, listed) if waited := time.Since(began); waited != 0 { t.Errorf("waited %s for the query, want no wait", waited) } synctest.Wait() wantQueries(t, dnsbl, 1, 0) waitForTheResolver() }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestVerdictUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) { synctest.Test(t, func(t *testing.T) { resolver := &resolverStandIn{answers: map[string]answer{ listedName: {addrs: []string{listing}}, }} dnsbl := newDNSBL(resolver, dnsblParams(zone)) wantZones(t, dnsbl, listed) wantZones(t, dnsbl, unlisted) synctest.Wait() // The zone lists the other client from now on, but the verdicts // kept are used, and the zone is not asked again, until the TTL // has passed. resolver.set(listedName, answer{rcode: noSuchName}) resolver.set(unlistedName, answer{addrs: []string{listing}}) time.Sleep(cacheTTL - time.Nanosecond) wantZones(t, dnsbl, listed, zone) wantZones(t, dnsbl, unlisted) synctest.Wait() wantQueries(t, dnsbl, 2, 0) // Then neither verdict is used, and both clients are asked about // again. time.Sleep(time.Nanosecond) wantZones(t, dnsbl, listed) wantZones(t, dnsbl, unlisted) synctest.Wait() wantQueries(t, dnsbl, 4, 0) wantZones(t, dnsbl, listed) wantZones(t, dnsbl, unlisted, zone) }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestQueryNotAnsweredWithinTheTimeoutFails(t *testing.T) { synctest.Test(t, func(t *testing.T) { queue := newQueue() p := dnsblParams(zone) p.Alerts = queue dnsbl := newDNSBL(&resolverStandIn{hanging: true}, p) wantZones(t, dnsbl, listed) time.Sleep(timeout - time.Nanosecond) synctest.Wait() wantQueries(t, dnsbl, 1, 0) time.Sleep(time.Nanosecond) synctest.Wait() wantQueries(t, dnsbl, 1, 1) if got := waiting(queue); len(got) != 1 || got[0].Detail["error"] != "ask the zone: i/o timeout" { t.Errorf("alerts waiting %+v, want the timeout's", got) } if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 { t.Errorf("verdicts %+v, want none", verdicts) } waitForTheResolver() }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestZoneThatFailsOrRefusesGivesNoVerdictAndIsLeftAloneForAMinute(t *testing.T) { for _, tc := range []struct { name string answer answer error string }{ { "a server failure", answer{rcode: serverFailure}, "ask the zone: server misbehaving", }, {"a refusal", answer{rcode: refused}, "ask the zone: server misbehaving"}, { "an answer in 127.255.255.0/24, with which Spamhaus refuses a query", answer{addrs: []string{"127.255.255.254"}}, "the zone refused the query: 127.255.255.254", }, { "an answer outside 127.0.0.0/8, as for a name that does not exist", answer{addrs: []string{"192.0.2.1"}}, "the answer is outside 127.0.0.0/8: 192.0.2.1", }, } { //nolint:paralleltest // one at a time, as the comment at the top of this file says t.Run(tc.name, func(t *testing.T) { synctest.Test(t, func(t *testing.T) { var log bytes.Buffer queue := newQueue() p := dnsblParams(zone) p.Alerts = queue p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil)) dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{ listedName: tc.answer, }}, p) // The failure gives no verdict, and the zone is not asked // again within a minute of it. wantZones(t, dnsbl, listed) synctest.Wait() time.Sleep(time.Minute - time.Nanosecond) wantZones(t, dnsbl, listed) synctest.Wait() wantQueries(t, dnsbl, 1, 1) time.Sleep(time.Nanosecond) wantZones(t, dnsbl, listed) synctest.Wait() wantQueries(t, dnsbl, 2, 2) if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 { t.Errorf("verdicts %+v, want none", verdicts) } // One alert for the first failure; the cooldown holds back // the second. wantAlert(t, queue, alerts.Alert{ Time: time.Now().Add(-time.Minute), Event: alerts.EventSourceFailure, Reason: "asking a DNSBL zone failed", Detail: map[string]any{"source": zone, "error": tc.error}, }) if !strings.Contains(log.String(), `"msg":"asking a DNSBL zone failed",`+ `"zone":"`+zone+`","error":"`+tc.error) { t.Errorf("logged\n%s\nwant the failures", log.String()) } }) }) } } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestAtMost1000QueriesUnderWay(t *testing.T) { synctest.Test(t, func(t *testing.T) { dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone)) client := netip.MustParseAddr("198.18.0.0") for range 1001 { dnsbl.ListedBy(t.Context(), client) client = client.Next() } synctest.Wait() wantQueries(t, dnsbl, 1000, 0) waitForTheResolver() }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestMetricsCountEachZonesQueriesAndThoseThatFailed(t *testing.T) { synctest.Test(t, func(t *testing.T) { resolver := &resolverStandIn{answers: map[string]answer{ listedName: {addrs: []string{listing}}, "99.2.0.192." + otherZone + ".": {rcode: serverFailure}, }} dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone)) m := metrics.New(1, "app") m.AddReputation(reputation.New(params()), dnsbl) wantZones(t, dnsbl, listed) synctest.Wait() scraped := httptest.NewRecorder() m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody)) for series, want := range map[string]string{ "queries_total" + `{instance="app",source="` + zone + `"}`: "1", "failures_total" + `{instance="app",source="` + zone + `"}`: "0", "queries_total" + `{instance="app",source="` + otherZone + `"}`: "1", "failures_total" + `{instance="app",source="` + otherZone + `"}`: "1", } { line := "\nsmallwebwaf_reputation_" + series + " " + want + "\n" if !strings.Contains(scraped.Body.String(), line) { t.Errorf("metrics\n%s\nwant%s", scraped.Body.String(), line) } } }) } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestVerdictsKeptAcrossARestart(t *testing.T) { synctest.Test(t, func(t *testing.T) { resolver := &resolverStandIn{answers: map[string]answer{ listedName: {addrs: []string{listing}}, }} dnsbl := newDNSBL(resolver, dnsblParams(zone)) fetched := time.Now() wantZones(t, dnsbl, unlisted) wantZones(t, dnsbl, listed) synctest.Wait() kept := dnsbl.Snapshot() want := []reputation.Verdict{ {Zone: zone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: fetched}, {Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: fetched}, } if !reflect.DeepEqual(kept, want) { t.Errorf("verdicts %+v, want %+v", kept, want) } // Restarted an hour later with what reputation.json keeps, it uses // the verdicts, and asks the zone nothing, until the TTL has passed // since they were fetched. time.Sleep(time.Hour) restarted := &resolverStandIn{} again := newDNSBL(restarted, dnsblParams(zone)) again.Load(kept) wantZones(t, again, listed, zone) wantZones(t, again, unlisted) synctest.Wait() wantAsked(t, restarted) time.Sleep(cacheTTL - time.Hour) wantZones(t, again, listed) synctest.Wait() wantAsked(t, restarted, listedName) }) } func TestNeitherAVerdictOfAZoneNotNamedNorOnePastItsTTLIsKept(t *testing.T) { t.Parallel() now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC) p := dnsblParams(zone) p.Now = func() time.Time { return now } dnsbl := reputation.NewDNSBL(p) // The last verdict still in use, one fetched a TTL ago, and one of a // zone SWWAF_DNSBL_ZONES does not name. inUse := reputation.Verdict{ Zone: zone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: now.Add(-cacheTTL + time.Nanosecond), } stale := reputation.Verdict{ Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: now.Add(-cacheTTL), } notNamed := reputation.Verdict{ Zone: otherZone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: now, } dnsbl.Load([]reputation.Verdict{notNamed, stale, inUse}) if got := dnsbl.Snapshot(); !reflect.DeepEqual(got, []reputation.Verdict{inUse}) { t.Errorf("verdicts %+v, want only %+v", got, inUse) } } func TestAtMost100000VerdictsKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) { t.Parallel() now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC) p := dnsblParams(zone) p.Now = func() time.Time { return now } dnsbl := reputation.NewDNSBL(p) // 100,001 verdicts, listed by client, as reputation.json lists them, // each fetched a millisecond before the one before it: the last is one // too many. const count = 100001 verdicts := make([]reputation.Verdict, 0, count) client := netip.MustParseAddr("198.18.0.0") for i := range count { verdicts = append(verdicts, reputation.Verdict{ Zone: zone, Client: client, Fetched: now.Add(-time.Duration(i) * time.Millisecond), }) client = client.Next() } dnsbl.Load(verdicts) got := dnsbl.Snapshot() if len(got) != count-1 || !slices.Contains(got, verdicts[0]) || slices.Contains(got, verdicts[count-1]) { t.Errorf("%d verdicts kept, want all but the one fetched longest ago", len(got)) } } //nolint:paralleltest // one at a time, as the comment at the top of this file says func TestQueriesGoToTheResolverSWWAFDNSBLResolverNames(t *testing.T) { resolver := &resolverStandIn{answers: map[string]answer{ listedName: {addrs: []string{listing}}, }} conn, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } served := make(chan struct{}) go func() { resolver.serveUDP(conn) close(served) }() t.Cleanup(func() { _ = conn.Close() <-served }) p := dnsblParams(zone) p.Resolver = netip.MustParseAddrPort(conn.LocalAddr().String()) // On the real clock: the stand-in answers at once, so only a test // process held up for a whole minute would see the query fail. p.Timeout = time.Minute isListed, err := reputation.NewDNSBL(p).LookUp(zone, netip.MustParseAddr(listed)) if err != nil || !isListed { t.Errorf("listed %t (%v), want true", isListed, err) } wantAsked(t, resolver, listedName) } // resolverStandIn is a stand-in for the resolver the zones are asked // through. It answers each query by the name asked about, as answers // gives, with no such name for a name answers does not give, and not at // all while hanging. It notes each name asked about. type resolverStandIn struct { mu sync.Mutex answers map[string]answer hanging bool names []string } // answer is how the stand-in answers a name: with an A record of each of // addrs, or with the response code rcode, unless it is 0, for no error. type answer struct { addrs []string rcode uint16 } // What the stand-in reads of a query, and writes in its reply. const ( // headerLength is the length of a DNS message's header, which the // question follows: its id, its flags, and how many questions, // answers and other records it holds, two bytes each. headerLength = 12 // typeAndClass is the length of the type and the class that end a // question, after its name. typeAndClass = 4 // replyFlags mark a reply to a query that asked for recursion, which // is available, with no error. The response code goes in their last // four bits. replyFlags = 0x8180 // maxMessage is the longest query read over UDP. maxMessage = 1232 ) // set has the stand-in answer name with given. func (s *resolverStandIn) set(name string, given answer) { s.mu.Lock() defer s.mu.Unlock() s.answers[name] = given } // dial connects Go's resolver to the stand-in through an in-memory // connection, on which it sends each query, and reads each reply, after // its length, as over TCP. func (s *resolverStandIn) dial(context.Context, string, string) (net.Conn, error) { client, server := net.Pipe() go s.serve(server) return client, nil } // serve answers the queries that come on conn until the resolver closes // it. func (s *resolverStandIn) serve(conn net.Conn) { defer func() { _ = conn.Close() }() for { var length [2]byte _, err := io.ReadFull(conn, length[:]) if err != nil { return } message := make([]byte, binary.BigEndian.Uint16(length[:])) _, err = io.ReadFull(conn, message) if err != nil { return } reply, answered := s.reply(message) if !answered { continue // the resolver gives up, and closes conn } //nolint:gosec // a reply of a few dozen bytes _, err = conn.Write(append(binary.BigEndian.AppendUint16(nil, uint16(len(reply))), reply...)) if err != nil { return } } } // serveUDP answers the queries that come on conn, each in a datagram, as // a resolver does, until conn is closed. func (s *resolverStandIn) serveUDP(conn net.PacketConn) { message := make([]byte, maxMessage) for { n, from, err := conn.ReadFrom(message) if err != nil { return } reply, answered := s.reply(message[:n]) if answered { _, _ = conn.WriteTo(reply, from) } } } // reply returns the stand-in's reply to message, a query, and false for // none, while it hangs. It notes the name asked about. func (s *resolverStandIn) reply(message []byte) ([]byte, bool) { // The name is labels, each after its length, ended by a length of 0. var labels []string end := headerLength for message[end] != 0 { length := int(message[end]) labels = append(labels, string(message[end+1:end+1+length])) end += 1 + length } end += 1 + typeAndClass name := strings.Join(labels, ".") + "." s.mu.Lock() s.names = append(s.names, name) given, found := s.answers[name] hanging := s.hanging s.mu.Unlock() if hanging { return nil, false } if !found { given = answer{rcode: noSuchName} } // The query's id, the flags, one question, the answers, and no other // records, then the question, as asked. reply := slices.Clone(message[:2]) reply = binary.BigEndian.AppendUint16(reply, replyFlags|given.rcode) reply = binary.BigEndian.AppendUint16(reply, 1) //nolint:gosec // a handful of answers reply = binary.BigEndian.AppendUint16(reply, uint16(len(given.addrs))) reply = append(reply, 0, 0, 0, 0) reply = append(reply, message[headerLength:end]...) // An A record starts with the name asked about, by a pointer to it in // the question, then its type, A, its class, IN, how long it may be // kept, 60 seconds, and the length of its address, 4 bytes. record := []byte{0xc0, headerLength, 0, 1, 0, 1, 0, 0, 0, 60, 0, 4} for _, addr := range given.addrs { reply = append(reply, record...) reply = append(reply, netip.MustParseAddr(addr).AsSlice()...) } return reply, true } // dnsblParams returns the DNSBLParams of zones, with the tests' cache TTL // and timeout, by the bubble's clock, with alerts to a queue that sends // none. func dnsblParams(zones ...string) reputation.DNSBLParams { return reputation.DNSBLParams{ Zones: zones, CacheTTL: cacheTTL, Timeout: timeout, Now: time.Now, ProcessLog: slog.New(slog.DiscardHandler), Alerts: newQueue(), } } // waitForTheResolver waits, on the bubble's clock, an hour, until Go's // resolver has given up on every stand-in that does not answer: it waits // for a server as long as /etc/resolv.conf has it wait, a few seconds, // even after the query was given up, and a bubble cannot end before it. func waitForTheResolver() { time.Sleep(time.Hour) } // newDNSBL returns the DNSBL of p, asking resolver. func newDNSBL(resolver *resolverStandIn, p reputation.DNSBLParams) *reputation.DNSBL { dnsbl := reputation.NewDNSBL(p) dnsbl.SetDial(resolver.dial) return dnsbl } // wantZones checks the zones whose verdict dnsbl says lists client, as a // request from client finds them. func wantZones(t *testing.T, dnsbl *reputation.DNSBL, client string, want ...string) { t.Helper() got := dnsbl.ListedBy(t.Context(), netip.MustParseAddr(client)) if !slices.Equal(got, want) { t.Errorf("%s is listed by %v, want %v", client, got, want) } } // wantQueries checks how many queries dnsbl made to zone, and how many of // them failed. func wantQueries(t *testing.T, dnsbl *reputation.DNSBL, queries, failures int) { t.Helper() if dnsbl.Queries(zone) != queries || dnsbl.Failures(zone) != failures { t.Errorf("%d queries and %d failures, want %d and %d", dnsbl.Queries(zone), dnsbl.Failures(zone), queries, failures) } } // wantAsked checks the names the stand-in was asked about, in any order. func wantAsked(t *testing.T, resolver *resolverStandIn, want ...string) { t.Helper() resolver.mu.Lock() got := slices.Sorted(slices.Values(resolver.names)) resolver.mu.Unlock() slices.Sort(want) if !slices.Equal(got, want) { t.Errorf("asked about %v, want %v", got, want) } }