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) } }