package proxy_test import ( "net" "net/http" "net/http/httptest" "net/netip" "slices" "sync" "testing" "testing/synctest" "time" "sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // asnAndCountry is what a lookup gives a client: its AS number, AS name // and country. type asnAndCountry struct{ asn, asName, country string } func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) { t.Parallel() // The stand-in for GeoJS answers only once released. A request that // waited for it would wait an hour, and get no answer within // waitLimit. geojsURL, asked, release := startHeldGeoJS(t) s, _, server := startWithClock(t, geojsURL, map[string]string{ lookupTimeout: "1h", rateLimitPerMinute: "1", }) // fromDE's second request breaks the limit and bans it, and fromKP // comes too. None waits for GeoJS. for _, line := range []logLine{ s.get(fromDE, http.StatusOK, requestlog.ActionForward), s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), s.get(fromKP, http.StatusOK, requestlog.ActionForward), } { got := asnAndCountry{line.ASN, line.ASName, line.Country} if got != (asnAndCountry{}) { t.Errorf("log line has %+v before GeoJS answered, want nothing", got) } } // Once GeoJS answers, each answer reaches the client's history, and // fromDE's reaches the notes of its ban. release() netblock := netip.MustParsePrefix(fromDE + "/32") waitUntil(func() bool { return historyOf(t, server, fromDE).ASN != "" && historyOf(t, server, fromKP).ASN != "" && server.Ledger.Bans(netblock)[0].Notes.ASN != "" }) de := asnAndCountry{asnDE, asNameDE, "DE"} for addr, want := range map[string]asnAndCountry{ fromDE: de, fromKP: {asnKP, asNameKP, "KP"}, } { h := historyOf(t, server, addr) if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want { t.Errorf("%s's history has %+v, want %+v", addr, got, want) } } notes := server.Ledger.Bans(netblock)[0].Notes if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de { t.Errorf("the ban's notes have %+v, want %+v", got, de) } // GeoJS was asked about each client once, fromKP after fromDE, whose // request was under way when fromKP came. if got := asked(); !slices.Equal(got, []string{fromDE, fromKP}) { t.Errorf("GeoJS was asked about %v, want %s and %s", got, fromDE, fromKP) } } func TestASNumberAndNameInTheLogLineTheHistoryTheBanNotesAndTheAlert(t *testing.T) { t.Parallel() geojsURL, _ := startGeoJS(t) app := startApp(t, func(http.ResponseWriter, *http.Request) {}) clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)} addr, out, server, queue := startProxyWithAlerts(t, app.URL, geojsURL, clk.Now, map[string]string{ trustedProxies: trustLocalhost, alertWebhookURL: "https://alerts.example/smallwebwaf", rateLimitPerMinute: "1", }) s := &sender{t: t, addr: addr, out: out} // The answer is kept before the requests, so GeoJS is not asked, and // gives no answer of its own. netblock := netip.MustParsePrefix(fromDE + "/32") server.GeoJS.Load([]lookup.Answer{{ Client: netblock, ASN: asnDE, ASName: asNameDE, Country: "DE", Answered: clk.Now(), Used: clk.Now(), }}) want := asnAndCountry{asnDE, asNameDE, "DE"} for _, line := range []logLine{ s.get(fromDE, http.StatusOK, requestlog.ActionForward), s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), } { if got := (asnAndCountry{line.ASN, line.ASName, line.Country}); got != want { t.Errorf("log line has %+v, want %+v", got, want) } } h := historyOf(t, server, fromDE) if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want || !h.LookedUp.Equal(clk.Now()) { t.Errorf("history has %+v, looked up at %s; want %+v, at %s", got, h.LookedUp, want, clk.Now()) } notes := server.Ledger.Bans(netblock)[0].Notes if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != want { t.Errorf("the ban's notes have %+v, want %+v", got, want) } waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] if len(waiting) != 1 { t.Fatalf("alerts waiting %+v, want the ban's alone", waiting) } alert := waiting[0] if got := (asnAndCountry{alert.ASN, alert.ASName, alert.Country}); got != want { t.Errorf("the ban's alert has %+v, want %+v", got, want) } } func TestLookupSourceOffLooksNoClientUp(t *testing.T) { t.Parallel() geojsURL, asked := startGeoJS(t) s, clk, server := startWithClock(t, geojsURL, map[string]string{lookupSource: "off"}) // Even an answer kept from before is not used. server.GeoJS.Load([]lookup.Answer{{ Client: netip.MustParsePrefix(fromDE + "/32"), ASN: asnDE, ASName: asNameDE, Country: "DE", Answered: clk.Now(), Used: clk.Now(), }}) for _, from := range []string{fromDE, fromKP} { line := s.get(from, http.StatusOK, requestlog.ActionForward) got := asnAndCountry{line.ASN, line.ASName, line.Country} if got != (asnAndCountry{}) { t.Errorf("log line has %+v, want nothing", got) } } if h := historyOf(t, server, fromDE); h.ASN != "" || !h.LookedUp.IsZero() { t.Errorf("history has %q, looked up at %s, want no lookup", h.ASN, h.LookedUp) } if len(asked()) != 0 { t.Errorf("GeoJS was asked about %v, want nothing", asked()) } } func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) { t.Parallel() var ( mu sync.Mutex got [][2][]string // each request's X-Client-ASN and X-Client-Country ) app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { mu.Lock() defer mu.Unlock() got = append(got, [2][]string{ r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country"), }) }) geojsURL, _ := startGeoJS(t) addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{ trustedProxies: trustLocalhost, addLookupHeaders: "true", }) s := &sender{t: t, addr: addr, out: out} // Each client sends headers of its own. fromDE's first request waits // for its answer, which the app is passed; unplaced has none to pass, // and a client on a private address is not looked up. for _, from := range []string{fromDE, unplaced, "10.0.0.8"} { s.requestWithHeader(from, "/", clientsOwnLookupHeaders, http.StatusOK, requestlog.ActionForward) } mu.Lock() defer mu.Unlock() want := [][2][]string{{{asnDE}, {"DE"}}, {nil, nil}, {nil, nil}} if !slices.EqualFunc(got, want, func(a, b [2][]string) bool { return slices.Equal(a[0], b[0]) && slices.Equal(a[1], b[1]) }) { t.Errorf("the app was passed %v, want %v", got, want) } } func TestClientsOwnLookupHeadersAreRemovedWhileTheSettingIsOff(t *testing.T) { t.Parallel() var ( mu sync.Mutex asn, country []string ) app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { mu.Lock() defer mu.Unlock() asn, country = r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country") }) geojsURL, _ := startGeoJS(t) addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{ trustedProxies: trustLocalhost, }) s := &sender{t: t, addr: addr, out: out} s.requestWithHeader(fromDE, "/", clientsOwnLookupHeaders, http.StatusOK, requestlog.ActionForward) mu.Lock() defer mu.Unlock() if asn != nil || country != nil { t.Errorf("the app was passed X-Client-ASN %v and X-Client-Country %v, want neither", asn, country) } } func TestRequestWaitsAsLongAsTheLookupTimeoutSays(t *testing.T) { t.Parallel() // The test runs in a synctest bubble, where the time package runs on a // clock of the test's own: the wait lasts exactly as long as it should, // however slowly the test process runs. Nothing in it may wait on the // network, which would keep that clock from moving on: the request is // handed to the proxy's handler, and GeoJS is one that never answers. synctest.Test(t, func(t *testing.T) { // Not the default second. The exclusive list needs the answer, and // the app is never reached. const timeout = 3 * time.Second server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL, time.Now, map[string]string{ lookupTimeout: timeout.String(), allowedCountries: "DE", }) req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody) req.RemoteAddr = net.JoinHostPort(fromDE, "1234") began := time.Now() server.Handler.ServeHTTP(httptest.NewRecorder(), req) if waited := time.Since(began); waited != timeout { t.Errorf("the request waited %s for its answer, want %s", waited, timeout) } // Without an answer, the client is in no country the list allows. wantLine(t, out.requestLine(t), http.StatusForbidden, requestlog.ActionCountryDenied) }) } // unansweredGeoJSURL is where a GeoJS that never answers is asked: a // request to it waits, without the network, until it is abandoned. // TestMain registers it with Go's default transport, through which GeoJS // is asked. const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json" func TestMain(m *testing.M) { transport, _ := http.DefaultTransport.(*http.Transport) transport.RegisterProtocol("unanswered", unansweredGeoJS{}) m.Run() } // unansweredGeoJS is the GeoJS at unansweredGeoJSURL. type unansweredGeoJS struct{} // RoundTrip waits until req is abandoned. func (unansweredGeoJS) RoundTrip(req *http.Request) (*http.Response, error) { <-req.Context().Done() return nil, req.Context().Err() } // clientsOwnLookupHeaders are the X-Client-ASN and X-Client-Country a // client sends of its own, each twice, in two cases. const clientsOwnLookupHeaders = "X-Client-ASN: AS1\r\nx-client-asn: AS2\r\n" + "X-CLIENT-COUNTRY: KP\r\nx-client-country: CN" // waitUntil waits until done reports true, for at most waitLimit. func waitUntil(done func() bool) { deadline := time.Now().Add(waitLimit) for !done() && time.Now().Before(deadline) { time.Sleep(pollInterval) } }