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) // The AS number GeoJS gives unplaced, 64512, counts as unknown. for i, sent := range []struct{ client, asn, asName, country string }{ {fromDE, asnDE, asNameDE, "DE"}, {fromKP, asnKP, asNameKP, "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.ASN != sent.asn || line.ASName != sent.asName || line.Country != sent.country { t.Errorf("log line has %q, %q and %q, want %q, %q and %q", line.ASN, line.ASName, line.Country, sent.asn, sent.asName, 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 := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "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 TestPrivateAddressIsNeverLookedUp(t *testing.T) { t.Parallel() for _, tc := range []struct { name string env map[string]string }{ {"no setting needs the lookup", nil}, {"a country list is set", map[string]string{deniedCountries: "kp"}}, } { 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) // "" sends no X-Forwarded-For: the client is 127.0.0.1. for i, sent := range []string{ "10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9", } { 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) for _, field := range []string{"asn", "as_name", "country"} { value, present := line.fields[field] if !present || value != "" { t.Errorf("log line for %q has %s %v, want an empty one", line.ClientIP, field, value) } } } // GeoJS is asked about up to 200 waiting clients at once, so once it // has been asked about fromDE, which comes last, it has been asked // about every client before it that waited for an answer. req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req.Header.Set(forwardedFor, fromDE) wantStatus(t, do(t, req), http.StatusOK) waitUntil(func() bool { return slices.Contains(asked(), fromDE) }) if got := asked(); !slices.Equal(got, []string{fromDE}) { t.Errorf("GeoJS was asked about %v, want %s alone", got, fromDE) } }) } } func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) { t.Parallel() for _, tc := range []struct { name string allowNets string status int action string }{ { "not in SWWAF_ALLOW_NETS", "", http.StatusForbidden, requestlog.ActionCountryDenied, }, { "in SWWAF_ALLOW_NETS", "10.0.0.7,fd00::/8", http.StatusOK, requestlog.ActionForward, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() app := startApp(t, func(http.ResponseWriter, *http.Request) {}) geojsURL, asked := startGeoJS(t) addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ trustedProxies: trustLocalhost, allowedCountries: "de", allowNets: tc.allowNets, }) wantAnswers(t, addr, out, []sentRequest{ {"10.0.0.7", tc.status, tc.action}, {"fd00::5", tc.status, tc.action}, }) 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, // each in an AS of its own, 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() geojsURL, asked, release := startHeldGeoJS(t) release() return geojsURL, asked } // startHeldGeoJS is startGeoJS for a stand-in that answers nothing until // release is called. Each request to it waits until then. func startHeldGeoJS(t *testing.T) (string, func() []string, func()) { t.Helper() var ( asked struct { mu sync.Mutex addrs []string } released = make(chan struct{}) once sync.Once ) 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() <-released answers := make([]geojsAnswer, 0, len(addrs)) for _, addr := range addrs { answers = append(answers, answerAbout(addr)) } err := json.NewEncoder(w).Encode(answers) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) } })) t.Cleanup(geojs.Close) release := func() { once.Do(func() { close(released) }) } // Run before geojs.Close, which waits for every request to be answered. t.Cleanup(release) return geojs.URL, func() []string { asked.mu.Lock() defer asked.mu.Unlock() return slices.Clone(asked.addrs) }, release } // The AS numbers and names the stand-in for GeoJS gives fromDE and // fromKP, as they are logged. const ( asnDE = "AS64496" asNameDE = "Example Net" asnKP = "AS64511" asNameKP = "Other Net" ) // geojsAnswer is an answer of GeoJS about one address, with the fields // smallwebwaf reads. // //nolint:tagliatelle // GeoJS's own names type geojsAnswer struct { IP string `json:"ip"` ASN int `json:"asn"` ASName string `json:"organization_name"` CountryCode string `json:"country_code,omitempty"` } // answerAbout is what the stand-in for GeoJS answers about addr: for an // address it cannot place, the AS number 64512 and the AS name Unknown // with no country, as GeoJS does. func answerAbout(addr string) geojsAnswer { switch addr { case fromDE: return geojsAnswer{IP: addr, ASN: 64496, ASName: asNameDE, CountryCode: "DE"} case fromKP: return geojsAnswer{IP: addr, ASN: 64511, ASName: asNameKP, CountryCode: "KP"} } return geojsAnswer{IP: addr, ASN: 64512, ASName: "Unknown"} }