Files
smallwebwaf/internal/lookup/lookup_test.go
T
clawbot 234c5eac60
check / check (push) Waiting to run
Serve Prometheus metrics behind SWWAF_METRICS_TOKEN (closes #23)
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. In the client's
history a 401 counts as refused, the metrics and the 404s as neither.
SWWAF_METRICS_TOP_N bounds the series by country, the rest counted as
other.

Deviation: go.mod and go.sum written by hand, 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
2026-10-06 11:40:27 +02:00

632 lines
16 KiB
Go

package lookup_test
import (
"encoding/json"
"log/slog"
"net/http"
"net/http/httptest"
"net/netip"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"github.com/prometheus/client_golang/prometheus/testutil"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics"
)
const (
// germany is where the stand-in for GeoJS places every address but
// unplaced.
germany = "DE"
// unplaced is the address it cannot place.
unplaced = "192.0.2.1"
// leftOut is the address it leaves out of its answer when
// answeringWithoutLeftOut.
leftOut = "203.0.113.7"
// timeout is how long a new client waits for its answer.
timeout = time.Second
// week is how long an answer is kept.
week = 7 * 24 * time.Hour
)
// The tests that have GeoJS asked run in a synctest bubble, where the time
// package runs on a clock of the test's own: a wait lasts exactly as long
// as it should, however slowly the test process runs, and synctest.Wait
// returns once g has done all it can before time passes. The stand-in for
// GeoJS answers without the network, since a request waiting on the
// network would keep that clock from moving on.
func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "")
wantRequests(t, geojs, 2)
// An answer without a country is kept too.
clock.advance(week - time.Second)
wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "")
wantRequests(t, geojs, 2)
clock.advance(time.Second)
wantCountry(t, g, placed, germany)
wantRequests(t, geojs, 3)
wantAsked(t, geojs, 2, "203.0.113.9")
})
}
func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
client := netip.MustParsePrefix("203.0.113.9/32")
// The client comes while GeoJS is asked about an earlier client, which
// it answers most of a second later. It is then asked about the client
// and does not answer: that request is abandoned a second after it
// began, well after the client's wait is over.
geojs.set(answeringSlowly)
var earlier sync.WaitGroup
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
defer earlier.Wait()
waitForRequests(t, geojs, 1)
geojs.set(hanging)
began := time.Now()
wantCountry(t, g, client, "")
took := time.Since(began)
if took != timeout {
t.Errorf("waited %s for the answer, want %s", took, timeout)
}
// Its next request does not wait.
began = time.Now()
wantCountry(t, g, client, "")
took = time.Since(began)
if took != 0 {
t.Errorf("waited %s again, want no wait", took)
}
// Once GeoJS answers, the client is asked about again in the
// background, and has its country.
geojs.set(answering)
waitForCountry(t, g, clock, client, germany)
})
}
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
answers int
// named is whether the answer names the other client asked about.
named bool
}{
{"null", answeringNull, false},
{"empty list", answeringEmptyList, false},
{"list without " + leftOut, answeringWithoutLeftOut, true},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
other := netip.MustParsePrefix("203.0.113.1/32")
client := netip.MustParsePrefix(leftOut + "/32")
// GeoJS fails, and is left alone for a second while the client
// comes too, so that the next request asks about both.
geojs.set(failing)
wantCountry(t, g, other, "")
wantCountry(t, g, client, "")
geojs.set(tc.answers)
clock.advance(time.Second)
wantCountry(t, g, other, "")
waitForRequests(t, geojs, 2)
// The answer counts as a failure, and the client is asked about
// again, with the other client only if the answer left it out too.
geojs.set(answering)
waitForCountry(t, g, clock, client, germany)
wantCountry(t, g, other, germany)
wantRequests(t, geojs, 3)
if tc.named {
wantAsked(t, geojs, 2, leftOut)
} else {
wantAsked(t, geojs, 2, leftOut, "203.0.113.1")
}
})
})
}
}
func TestRedirectCountsAsFailure(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, _, g := start()
geojs.set(redirecting)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
wantRequests(t, geojs, 1)
})
}
func TestCountryIsKeptInCapitals(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, _, g := start()
geojs.set(answeringInLowerCase)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany)
})
}
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log strings.Builder
// GeoJS does not answer, so the request to it is abandoned, and fails.
geojs := &standIn{answers: hanging}
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1),
})
g.SetTransport(geojs)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
synctest.Wait()
logged := log.String()
if !strings.Contains(logged, "asking GeoJS failed") ||
strings.Contains(logged, "203.0.113.9") {
t.Errorf("logged %q, want the failure without the address asked about", logged)
}
})
}
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
// GeoJS fails, and is then left alone for a second, while three more
// clients come. An IPv6 client is a /64, and GeoJS is asked about its
// first address.
geojs.set(failing)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "")
wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "")
wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "")
wantRequests(t, geojs, 1)
geojs.set(answering)
clock.advance(time.Second)
wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany)
wantRequests(t, geojs, 2)
wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3")
})
}
func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
clients := newClients()
kept := clients()
wantCountry(t, g, kept, germany)
geojs.set(failing)
wantCountry(t, g, kept, germany)
wantRequests(t, geojs, 1)
// Each failure leaves GeoJS alone twice as long as the one before, up
// to five minutes. New clients meanwhile count as not found, and the
// client with a kept answer still gets its country, without GeoJS being
// asked.
requests := 1
for _, delay := range []time.Duration{
time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second,
16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second,
256 * time.Second, 5 * time.Minute, 5 * time.Minute,
} {
wantCountry(t, g, clients(), "")
requests++
wantRequests(t, geojs, requests)
clock.advance(delay - time.Millisecond)
wantCountry(t, g, clients(), "")
wantCountry(t, g, kept, germany)
wantRequests(t, geojs, requests)
clock.advance(time.Millisecond)
}
// Once GeoJS answers again, it is asked about every client waiting.
geojs.set(answering)
wantCountry(t, g, clients(), germany)
wantRequests(t, geojs, requests+1)
asked := waitForRequests(t, geojs, requests+1)
if len(asked[requests]) != 23 {
t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests]))
}
})
}
func TestAtMost200AddressesInOneRequest(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
clients := newClients()
first := clients()
// 201 clients wait while GeoJS is left alone after a failure.
geojs.set(failing)
wantCountry(t, g, first, "")
for range 200 {
wantCountry(t, g, clients(), "")
}
// The first one's next request has GeoJS asked again.
geojs.set(answering)
clock.advance(time.Second)
wantCountry(t, g, first, "")
asked := waitForRequests(t, geojs, 3)
if len(asked[1]) != 200 || len(asked[2]) != 1 {
t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1",
len(asked[1]), len(asked[2]))
}
})
}
func TestAtMost10000ClientsWait(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g := start()
clients := newClients()
first := clients()
// 10,000 clients wait while GeoJS is left alone after a failure, and
// one more cannot join them.
geojs.set(failing)
wantCountry(t, g, first, "")
for range 9999 {
wantCountry(t, g, clients(), "")
}
extra := clients()
wantCountry(t, g, extra, "")
// The first one's next request has GeoJS asked about the 10,000, 200
// at a time, and not about the one more.
geojs.set(answering)
clock.advance(time.Second)
wantCountry(t, g, first, "")
asked := waitForRequests(t, geojs, 51)
for i, request := range asked {
if slices.Contains(request, extra.Addr().String()) {
t.Errorf("request %d asked about %s", i, extra.Addr())
}
}
// With room among those waiting, it is asked about.
wantCountry(t, g, extra, germany)
})
}
func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
m := metrics.New(1)
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
})
g.SetTransport(&standIn{answers: failing})
clients := newClients()
// GeoJS fails, so the first client goes without an answer, and GeoJS
// is left alone for a second, which does not pass in this test.
wantCountry(t, g, clients(), "")
wantUnanswered(t, m, 1)
// Meanwhile each new client goes without one at once, while there is
// room for it among the 10,000 that may wait.
for range 9999 {
wantCountry(t, g, clients(), "")
}
wantUnanswered(t, m, 10000)
// One more, for which there is no room, goes without one too.
wantCountry(t, g, clients(), "")
wantUnanswered(t, m, 10001)
})
}
// How the stand-in for GeoJS answers.
const (
answering = iota
answeringSlowly // most of a second later
answeringInLowerCase // with each country in lower case
answeringWithoutLeftOut // with a list that leaves leftOut out
answeringEmptyList // with []
answeringNull // with null
failing // with 503
hanging // not at all, until the request is abandoned
redirecting // with a redirect to itself
)
// standIn is a stand-in for GeoJS. It notes the addresses each request
// asks about.
type standIn struct {
mu sync.Mutex
answers int
requests [][]string
}
// RoundTrip has the stand-in answer req, in place of the network. A request
// abandoned before the stand-in answers fails, as over the network.
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
answer := httptest.NewRecorder()
s.ServeHTTP(answer, req)
err := req.Context().Err()
if err != nil {
return nil, err
}
return answer.Result(), nil
}
// ServeHTTP answers a request about the addresses in its ip parameter.
func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
addrs := strings.Split(r.URL.Query().Get("ip"), ",")
s.mu.Lock()
s.requests = append(s.requests, addrs)
answers := s.answers
s.mu.Unlock()
switch answers {
case failing:
w.WriteHeader(http.StatusServiceUnavailable)
return
case hanging:
<-r.Context().Done()
return
case redirecting:
http.Redirect(w, r, "/", http.StatusFound)
return
case answeringSlowly:
select {
case <-time.After(timeout * 4 / 5):
case <-r.Context().Done():
return
}
}
list := make([]map[string]string, 0, len(addrs))
for _, addr := range addrs {
country := germany
switch {
case addr == unplaced:
country = ""
case addr == leftOut && answers == answeringWithoutLeftOut:
continue
case answers == answeringInLowerCase:
country = strings.ToLower(germany)
}
list = append(list, map[string]string{"ip": addr, "country": country})
}
var answer any = list
switch answers {
case answeringEmptyList:
answer = []string{}
case answeringNull:
answer = nil
}
err := json.NewEncoder(w).Encode(answer)
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
}
}
// set sets how the stand-in answers.
func (s *standIn) set(answers int) {
s.mu.Lock()
defer s.mu.Unlock()
s.answers = answers
}
// asked returns the addresses each request has asked about so far.
func (s *standIn) asked() [][]string {
s.mu.Lock()
defer s.mu.Unlock()
return slices.Clone(s.requests)
}
// testClock is a clock the test sets. GeoJS tells the time by it, while
// waits run on the bubble's clock.
type testClock struct {
mu sync.Mutex
now time.Time
}
// Now tells the time.
func (c *testClock) Now() time.Time {
c.mu.Lock()
defer c.mu.Unlock()
return c.now
}
// advance moves the clock on by d.
func (c *testClock) advance(d time.Duration) {
c.mu.Lock()
defer c.mu.Unlock()
c.now = c.now.Add(d)
}
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
// asking the stand-in by that clock.
func start() (*standIn, *testClock, *lookup.GeoJS) {
geojs := &standIn{}
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
g := lookup.New(lookup.Params{
URL: lookup.URL,
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1),
})
g.SetTransport(geojs)
return geojs, clock, g
}
// newClients returns what returns a new IPv4 client each time it is
// called.
func newClients() func() netip.Prefix {
addr := netip.MustParseAddr("10.0.0.0")
return func() netip.Prefix {
addr = addr.Next()
return netip.PrefixFrom(addr, addr.BitLen())
}
}
// wantCountry checks the country g gives client.
func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) {
t.Helper()
got := g.Country(t.Context(), client)
if got != want {
t.Errorf("%s is in %q, want %q", client, got, want)
}
}
// wantRequests checks how many requests GeoJS has had.
func wantRequests(t *testing.T, geojs *standIn, want int) {
t.Helper()
got := len(geojs.asked())
if got != want {
t.Errorf("GeoJS had %d requests, want %d", got, want)
}
}
// wantAsked checks the addresses request i asked about, in any order.
func wantAsked(t *testing.T, geojs *standIn, i int, want ...string) {
t.Helper()
asked := geojs.asked()
if len(asked) <= i {
t.Fatalf("GeoJS had %d requests, want more than %d", len(asked), i)
}
got := slices.Sorted(slices.Values(asked[i]))
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("request %d asked about %v, want %v", i, got, want)
}
}
// wantUnanswered checks how many requests m counts as having gone without
// an answer from GeoJS.
func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
t.Helper()
got := testutil.ToFloat64(m.GeoJSUnanswered)
if got != want {
t.Errorf("%v requests went without an answer, want %v", got, want)
}
}
// waitForRequests waits until g has done all it can before time passes,
// checks that GeoJS has had count requests, and returns the addresses each
// asked about.
func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string {
t.Helper()
synctest.Wait()
asked := geojs.asked()
if len(asked) != count {
t.Fatalf("GeoJS had %d requests, want %d", len(asked), count)
}
return asked
}
// waitForCountry lets a request to GeoJS under way be abandoned, and moves
// the clock on a minute, so that GeoJS may be asked again after a failure.
// It then checks that client's next request does not wait but has it asked
// about again in the background, after which g gives it the country want.
func waitForCountry(
t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string,
) {
t.Helper()
time.Sleep(timeout)
clock.advance(time.Minute)
wantCountry(t, g, client, "")
synctest.Wait()
wantCountry(t, g, client, want)
}