Files
smallwebwaf/internal/proxy/countries_test.go
T
clawbot 9c67b5b837
check / check (push) Successful in 2m18s
Country allow and deny lists, looked up through GeoJS (closes #44)
SWWAF_DENIED_COUNTRIES and SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES refuse a
request with 403 before its body is read and before the rate limits count
it, logged as country_denied; every log line gains country. The new
internal/lookup asks GeoJS only while a list is set, one request at a
time carrying up to 200 waiting clients, keeps answers 7 days (at most
100,000), and after a failure waits a second, doubling to five minutes.
Private, loopback and link-local clients have no country and are never
sent. Codes are checked with golang.org/x/text/language.

Deviation from SPEC.md, per the issue: no SWWAF_LOOKUP_SOURCE or SWWAF_LOOKUP_TIMEOUT; 403, not SWWAF_BAN_RESPONSE.
Deviation: GeoJS's country endpoint, not geo.json, since only the country is needed.
Judgement call: an IPv6 /64 is asked about by its first address; at most 10,000 clients wait.
Deviation: go.mod and go.sum hand-written; no make target tidies them.

Model: opus-5-5
2026-10-04 05:08:54 +00:00

214 lines
5.6 KiB
Go

package proxy_test
import (
"encoding/json"
"maps"
"net/http"
"net/http/httptest"
"slices"
"strings"
"sync"
"sync/atomic"
"testing"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The clients the stand-in for GeoJS knows about.
const (
// fromDE is placed in Germany.
fromDE = client
// fromKP is placed in North Korea.
fromKP = "198.51.100.7"
// unplaced cannot be placed in any country.
unplaced = "192.0.2.1"
)
func TestCountryLists(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
refused []string
}{
{"denied", map[string]string{deniedCountries: "kp"}, []string{fromKP}},
{
"exclusively allowed", map[string]string{allowedCountries: "DE"},
[]string{fromKP, unplaced},
},
{
"both", map[string]string{deniedCountries: "kp", allowedCountries: "de,fr"},
[]string{fromKP, unplaced},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
geojsURL, _ := startGeoJS(t)
env := map[string]string{trustedProxies: trustLocalhost}
maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
for i, sent := range []struct{ client, country string }{
{fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""},
} {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, sent.client)
got := do(t, req)
line := out.requestLines(t, i+1)[i]
if line.Country != sent.country {
t.Errorf("log line has country %q, want %q", line.Country, sent.country)
}
if slices.Contains(tc.refused, sent.client) {
wantStatus(t, got, http.StatusForbidden)
wantLine(t, line, http.StatusForbidden, requestlog.ActionCountryDenied)
} else {
wantStatus(t, got, http.StatusOK)
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
}
}
if int(calls.Load()) != 3-len(tc.refused) {
t.Errorf("the app was called %d times, want %d",
calls.Load(), 3-len(tc.refused))
}
})
}
}
func TestCountryRefusalComesBeforeTheBodyAndTheRateLimits(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
geojsURL, _ := startGeoJS(t)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{
trustedProxies: trustLocalhost,
deniedCountries: "kp",
rateLimitPerMinute: "1",
})
// With a limit of one request a minute, the second request would be
// over it, were refused requests counted.
for i := range 2 {
req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("a body"))
req.Header.Set(forwardedFor, fromKP)
wantStatus(t, do(t, req), http.StatusForbidden)
line := out.requestLines(t, i+1)[i]
wantLine(t, line, http.StatusForbidden, requestlog.ActionCountryDenied)
if line.RequestBytes != 0 || line.LimitHit != "" {
t.Errorf("log line has request_bytes %d and limit_hit %q, want 0 and none",
line.RequestBytes, line.LimitHit)
}
}
if calls.Load() != 0 {
t.Errorf("the app was called %d times, want none", calls.Load())
}
}
func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
clients []string // "" sends no X-Forwarded-For: the client is 127.0.0.1
}{
{"no country list is set", nil, []string{fromKP, fromDE}},
{
"private, loopback and link-local addresses",
map[string]string{allowedCountries: "de"},
[]string{"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9"},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
geojsURL, asked := startGeoJS(t)
env := map[string]string{trustedProxies: trustLocalhost}
maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
for i, sent := range tc.clients {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
if sent != "" {
req.Header.Set(forwardedFor, sent)
}
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLines(t, i+1)[i]
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
country, present := line.fields["country"]
if !present || country != "" {
t.Errorf("log line for %q has country %v, want an empty one",
line.ClientIP, country)
}
}
if len(asked()) != 0 {
t.Errorf("GeoJS was asked about %v, want nothing", asked())
}
})
}
}
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP
// and no other address. It returns its URL, and what returns the
// addresses it has been asked about.
func startGeoJS(t *testing.T) (string, func() []string) {
t.Helper()
places := map[string]string{fromDE: "DE", fromKP: "KP"}
var asked struct {
mu sync.Mutex
addrs []string
}
geojs := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
addrs := strings.Split(r.URL.Query().Get("ip"), ",")
asked.mu.Lock()
asked.addrs = append(asked.addrs, addrs...)
asked.mu.Unlock()
answers := make([]map[string]string, 0, len(addrs))
for _, addr := range addrs {
answers = append(answers, map[string]string{
"ip": addr, "country": places[addr],
})
}
err := json.NewEncoder(w).Encode(answers)
if err != nil {
http.Error(w, err.Error(), http.StatusInternalServerError)
}
}))
t.Cleanup(geojs.Close)
return geojs.URL, func() []string {
asked.mu.Lock()
defer asked.mu.Unlock()
return slices.Clone(asked.addrs)
}
}