Country allow and deny lists, looked up through GeoJS (closes #44)
check / check (push) Failing after 3s
check / check (push) Failing after 3s
SWWAF_DENIED_COUNTRIES and SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES refuse a request with 403 before its body is read or rate-limited, logged as country_denied. internal/lookup asks GeoJS only while a list is set, 200 clients per request, one at a time, keeping answers 7 days. Failures, a redirect or an answer leaving an address out included, are logged without addresses; GeoJS is then left alone a second, doubling to five minutes. Private, loopback and link-local clients are never sent. Deviation, per the issue: no SWWAF_LOOKUP_SOURCE or SWWAF_LOOKUP_TIMEOUT; 403, not SWWAF_BAN_RESPONSE. Deviation: GeoJS's country endpoint, not geo.json. Judgement call: an IPv6 /64 is asked about by its first address; at most 10,000 clients wait. Judgement call: config.go lists the ISO 3166-1 codes; no widely used library holds them. Model: opus-5-5
This commit was merged in pull request #54.
This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"slices"
|
||||
)
|
||||
|
||||
// countryDenied reports whether the country lists refuse the request.
|
||||
// The client's country is looked up only while a list is set, and never
|
||||
// for a client on a private, loopback or link-local address, which has
|
||||
// no country and which neither list checks. A client whose country
|
||||
// cannot be found is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES.
|
||||
// ctx is the request's own context.
|
||||
func (rq *request) countryDenied(ctx context.Context) bool {
|
||||
denied := rq.h.config.DeniedCountries
|
||||
allowed := rq.h.config.ExclusivelyAllowedCountries
|
||||
|
||||
if len(denied) == 0 && len(allowed) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
if !hasCountry(rq.client) {
|
||||
return false
|
||||
}
|
||||
|
||||
country := rq.h.geojs.Country(ctx, clientGroup(rq.client))
|
||||
rq.line.Country = country
|
||||
|
||||
if slices.Contains(denied, country) {
|
||||
return true
|
||||
}
|
||||
|
||||
return len(allowed) > 0 && !slices.Contains(allowed, country)
|
||||
}
|
||||
|
||||
// hasCountry reports whether addr can be placed in a country: private,
|
||||
// loopback and link-local addresses cannot.
|
||||
func hasCountry(addr netip.Addr) bool {
|
||||
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
package proxy_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"maps"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"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 TestCountryRefusalComesBeforeTheBody(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",
|
||||
})
|
||||
|
||||
req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("a body"))
|
||||
req.Header.Set(forwardedFor, fromKP)
|
||||
wantStatus(t, do(t, req), http.StatusForbidden)
|
||||
|
||||
line := out.requestLine(t)
|
||||
wantLine(t, line, http.StatusForbidden, requestlog.ActionCountryDenied)
|
||||
|
||||
if line.RequestBytes != 0 {
|
||||
t.Errorf("log line has request_bytes %d, want 0", line.RequestBytes)
|
||||
}
|
||||
|
||||
if calls.Load() != 0 {
|
||||
t.Errorf("the app was called %d times, want none", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// The stand-in for GeoJS fails until placing is set, and then places
|
||||
// every address in Germany.
|
||||
var placing atomic.Bool
|
||||
|
||||
geojs := httptest.NewServer(http.HandlerFunc(
|
||||
func(w http.ResponseWriter, r *http.Request) {
|
||||
if !placing.Load() {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}}
|
||||
|
||||
err := json.NewEncoder(w).Encode(answer)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(geojs.Close)
|
||||
|
||||
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||
addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{
|
||||
trustedProxies: trustLocalhost,
|
||||
allowedCountries: "de",
|
||||
rateLimitPerMinute: "1",
|
||||
})
|
||||
|
||||
request := func() answer {
|
||||
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
||||
req.Header.Set(forwardedFor, fromDE)
|
||||
|
||||
return do(t, req)
|
||||
}
|
||||
|
||||
// While GeoJS fails, the client's country cannot be found, and its
|
||||
// request is refused.
|
||||
wantStatus(t, request(), http.StatusForbidden)
|
||||
|
||||
// Once GeoJS places it, a second after the failure, its requests are let
|
||||
// through. No refused one was counted, so the first let through is
|
||||
// within the limit of one a minute.
|
||||
placing.Store(true)
|
||||
|
||||
deadline := time.Now().Add(waitLimit)
|
||||
got := request()
|
||||
|
||||
for got.status == http.StatusForbidden && time.Now().Before(deadline) {
|
||||
time.Sleep(pollInterval)
|
||||
|
||||
got = request()
|
||||
}
|
||||
|
||||
wantStatus(t, got, http.StatusOK)
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
+11
-1
@@ -11,6 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
)
|
||||
|
||||
@@ -40,6 +41,9 @@ type Params struct {
|
||||
RequestLog io.Writer
|
||||
// ProcessLog receives the process's own messages.
|
||||
ProcessLog *slog.Logger
|
||||
// GeoJSURL is where clients' countries are looked up, normally
|
||||
// lookup.URL. GeoJS is asked only while a country list is set.
|
||||
GeoJSURL string
|
||||
}
|
||||
|
||||
// New returns the server smallwebwaf runs: each request it reads passes
|
||||
@@ -63,6 +67,11 @@ func New(params Params) *http.Server {
|
||||
PerHour: params.Config.RateLimitPerHour,
|
||||
PerDay: params.Config.RateLimitPerDay,
|
||||
}),
|
||||
geojs: lookup.New(lookup.Params{
|
||||
URL: params.GeoJSURL,
|
||||
Now: time.Now,
|
||||
ProcessLog: params.ProcessLog,
|
||||
}),
|
||||
},
|
||||
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
||||
IdleTimeout: clientIdleTimeout,
|
||||
@@ -80,6 +89,7 @@ type handler struct {
|
||||
errorLog *log.Logger
|
||||
transport http.RoundTripper
|
||||
limiter *ratelimit.Limiter
|
||||
geojs *lookup.GeoJS
|
||||
}
|
||||
|
||||
// newTransport returns what carries requests to the app. It never goes
|
||||
@@ -101,7 +111,7 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
rq := h.newRequest(w, r)
|
||||
defer rq.finish()
|
||||
|
||||
refused := rq.check()
|
||||
refused := rq.check(r.Context())
|
||||
if refused != nil {
|
||||
rq.answer(*refused)
|
||||
|
||||
|
||||
@@ -49,6 +49,8 @@ const (
|
||||
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
|
||||
)
|
||||
|
||||
// output collects what smallwebwaf writes on stdout.
|
||||
@@ -163,6 +165,16 @@ func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server {
|
||||
func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) {
|
||||
t.Helper()
|
||||
|
||||
return startProxyWithGeoJS(t, appURL, "", env)
|
||||
}
|
||||
|
||||
// startProxyWithGeoJS is startProxy with clients' countries looked up at
|
||||
// geojsURL.
|
||||
func startProxyWithGeoJS(
|
||||
t *testing.T, appURL, geojsURL string, env map[string]string,
|
||||
) (string, *output) {
|
||||
t.Helper()
|
||||
|
||||
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
|
||||
maps.Copy(settings, env)
|
||||
|
||||
@@ -180,6 +192,7 @@ func startProxy(t *testing.T, appURL string, env map[string]string) (string, *ou
|
||||
Config: cfg,
|
||||
RequestLog: out,
|
||||
ProcessLog: requestlog.NewProcessLogger(out),
|
||||
GeoJSURL: geojsURL,
|
||||
})
|
||||
|
||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||
|
||||
@@ -103,9 +103,18 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
||||
|
||||
// check is the one place where a request can be refused once its client
|
||||
// is known, before its body is read or anything reaches the app. It
|
||||
// returns nil to let the request through. The rate limits come first, so
|
||||
// that every request is counted, one refused for its size too.
|
||||
func (rq *request) check() *refusal {
|
||||
// returns nil to let the request through. The country lists come first,
|
||||
// and a request they refuse is not counted for the rate limits; then the
|
||||
// rate limits, so that every other request is counted, one refused for
|
||||
// its size too. ctx is the request's own context.
|
||||
func (rq *request) check(ctx context.Context) *refusal {
|
||||
if rq.countryDenied(ctx) {
|
||||
return &refusal{
|
||||
status: http.StatusForbidden,
|
||||
action: requestlog.ActionCountryDenied,
|
||||
}
|
||||
}
|
||||
|
||||
limitHit := rq.h.limiter.Count(clientGroup(rq.client), rq.start)
|
||||
if limitHit != "" {
|
||||
rq.line.LimitHit = limitHit
|
||||
|
||||
Reference in New Issue
Block a user