package proxy_test import ( "bufio" "errors" "io" "maps" "net/http" "net/netip" "slices" "sync" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) const ( // otherClient is a client next to client. otherClient = "203.0.113.10" // userAgent is the user agent of every request a sender sends. userAgent = "ban-test/1.0" // permanent is the log line's ban_expires for a permanent ban. permanent = "permanent" ) func TestBrokenLimitBansTheClient(t *testing.T) { t.Parallel() s, clk, _ := startWithClock(t, "", map[string]string{rateLimitPerMinute: "1"}) expires := requestlog.FormatTime(clk.Now().Add(time.Hour)) // The request over the limit of one a minute is refused, and bans the // client for an hour, the default. s.get(client, http.StatusOK, requestlog.ActionForward) line := s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit || line.BanExpires != expires { t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+ "want minute, limit and %s", line.LimitHit, line.Offence, line.BanExpires, expires) } // Every request while the ban lasts is refused. clk.advance(time.Hour - time.Second) line = s.get(client, http.StatusForbidden, requestlog.ActionBanned) if line.BanExpires != expires || line.Offence != "" || line.LimitHit != "" { t.Errorf("log line has ban_expires %q, offence %q and limit_hit %q, "+ "want %s and neither of the others", line.BanExpires, line.Offence, line.LimitHit, expires) } // Once it ends, the client is let through. clk.advance(time.Second) s.get(client, http.StatusOK, requestlog.ActionForward) } func TestBanLengthsFollowTheSettings(t *testing.T) { t.Parallel() s, clk, _ := startWithClock(t, "", map[string]string{ rateLimitPerMinute: "1", limitBanDuration: "10m", limitBanRepeatWindow: "1h", maxBanDuration: "1h", }) // breakLimit has client go over the limit of one a minute, and // returns when the ban that makes ends. breakLimit := func() string { s.get(client, http.StatusOK, requestlog.ActionForward) return s.get(client, http.StatusForbidden, requestlog.ActionRateLimited).BanExpires } wantExpires := func(got string, length time.Duration) { t.Helper() want := requestlog.FormatTime(clk.Now().Add(length)) if got != want { t.Errorf("ban ends at %s, want %s", got, want) } } // A first ban lasts SWWAF_LIMIT_BAN_DURATION; one within // SWWAF_LIMIT_BAN_REPEAT_WINDOW after it ended, three times as long. wantExpires(breakLimit(), 10*time.Minute) clk.advance(10*time.Minute + time.Hour) wantExpires(breakLimit(), 30*time.Minute) // Later than that, SWWAF_LIMIT_BAN_DURATION again. clk.advance(30*time.Minute + time.Hour + time.Second) wantExpires(breakLimit(), 10*time.Minute) clk.advance(10 * time.Minute) wantExpires(breakLimit(), 30*time.Minute) // 90 minutes would be longer than SWWAF_MAX_BAN_DURATION: the ban is // permanent. clk.advance(30 * time.Minute) got := breakLimit() if got != permanent { t.Errorf("ban ends at %s, want a permanent one", got) } clk.advance(365 * 24 * time.Hour) s.get(client, http.StatusForbidden, requestlog.ActionBanned) } func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) { t.Parallel() s, clk, server := startWithClock(t, "", map[string]string{rateLimitPerDay: "2"}) // The third request in a day is over the limit of two, and bans the // client for an hour. s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) for range 3 { s.get(client, http.StatusForbidden, requestlog.ActionBanned) } // Later the same day the client has its whole allowance again: the // ban set its counters back to zero, and the requests it refused were // not counted for the rate limits, only in its notes. clk.advance(time.Hour) s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32")) if len(banned) != 2 || banned[0].Notes.Refused != 3 { t.Errorf("bans %+v, want two, the first with 3 requests refused", banned) } } func TestBanCoversTheClientsNetblock(t *testing.T) { t.Parallel() // In the IPv4 cases, client breaks the limit; these two are next to it. const ( allowed = "203.0.113.60" // in SWWAF_ALLOW_NETS exempt = "203.0.113.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS ) for _, tc := range []struct { name string env map[string]string breaker string // the client that breaks the limit refused []string let []string // let through }{ { "an IPv4 address, by default", nil, client, nil, []string{otherClient, exempt}, }, { "the IPv4 netblock SWWAF_BAN_SCOPE_V4_PREFIX sets", map[string]string{banScopeV4Prefix: "24"}, client, []string{otherClient, exempt}, []string{"203.0.112.9", allowed}, }, { "an IPv6 /64", nil, "2001:db8:5::1", []string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"}, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() env := map[string]string{ rateLimitPerMinute: "1", allowNets: allowed, rateLimitExemptNets: exempt, } maps.Copy(env, tc.env) s, _, _ := startWithClock(t, "", env) s.get(tc.breaker, http.StatusOK, requestlog.ActionForward) s.get(tc.breaker, http.StatusForbidden, requestlog.ActionRateLimited) for _, sent := range tc.refused { s.get(sent, http.StatusForbidden, requestlog.ActionBanned) } for _, sent := range tc.let { s.get(sent, http.StatusOK, requestlog.ActionForward) } }) } } func TestBannedClientIsRefusedBeforeItsCountryIsLookedUp(t *testing.T) { t.Parallel() geojsURL, asked := startGeoJS(t) s, _, _ := startWithClock(t, geojsURL, map[string]string{ rateLimitPerMinute: "1", banScopeV4Prefix: "24", deniedCountries: "kp", }) // fromDE's ban covers otherClient, which is refused unasked about. s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) line := s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned) if line.Country != "" { t.Errorf("log line has country %q, want none", line.Country) } if !slices.Equal(asked(), []string{fromDE}) { t.Errorf("GeoJS was asked about %v, want %s alone", asked(), fromDE) } } func TestBanResponseAnswersEveryRefusalButTheSizeLimits(t *testing.T) { t.Parallel() const denied = "192.0.2.50" // in SWWAF_DENY_NETS for _, tc := range []struct { setting string // "" leaves SWWAF_BAN_RESPONSE at its default status int // 0 is the connection closed without an answer }{ {"", http.StatusForbidden}, {"403", http.StatusForbidden}, {"429", http.StatusTooManyRequests}, {"close", 0}, } { t.Run(banResponse+"="+tc.setting, func(t *testing.T) { t.Parallel() geojsURL, _ := startGeoJS(t) env := map[string]string{ rateLimitPerMinute: "1", denyNets: denied, deniedCountries: "kp", } if tc.setting != "" { env[banResponse] = tc.setting } s, _, _ := startWithClock(t, geojsURL, env) s.get(denied, tc.status, requestlog.ActionDenied) s.get(fromKP, tc.status, requestlog.ActionCountryDenied) s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, tc.status, requestlog.ActionRateLimited) s.get(fromDE, tc.status, requestlog.ActionBanned) }) } } func TestBanNotes(t *testing.T) { t.Parallel() geojsURL, _ := startGeoJS(t) s, clk, server := startWithClock(t, geojsURL, map[string]string{ rateLimitPerMinute: "1", deniedCountries: "kp", }) start := clk.Now() s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.request(fromDE, "/repo/commits?page=2", http.StatusForbidden, requestlog.ActionRateLimited) s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned) s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned) netblock := netip.MustParsePrefix(fromDE + "/32") want := bans.Ban{ Netblock: netblock, Start: start, Expires: start.Add(time.Hour), Cause: bans.CauseLimit, Notes: bans.Notes{ Country: "DE", Limit: 1, Window: minute, Count: 2, Request: bans.Request{ Time: start, Method: http.MethodGet, Host: appHost, Path: "/repo/commits?page=2", Status: http.StatusForbidden, UserAgent: userAgent, }, // The one let through, the one that broke the limit and the two // refused under the ban. Requests: 4, Refused: 2, EarlierBans: bans.EarlierBans{}, }, } ledger := server.Ledger got := ledger.Bans(netblock) if len(got) != 1 || got[0] != want { t.Fatalf("bans\n%+v\nwant\n%+v", got, want) } // The next ban counts this one among the earlier. clk.advance(time.Hour) s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) got = ledger.Bans(netblock) if len(got) != 2 || got[1].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { t.Errorf("bans %+v, want two, the second with one earlier ban for a limit", got) } } func TestMaxBansDropsTheBanOfTheNetblockSeenLongestAgo(t *testing.T) { t.Parallel() s, _, _ := startWithClock(t, "", map[string]string{ rateLimitPerMinute: "1", maxBans: "1", }) // One ban is held, so otherClient's ban drops client's. s.get(client, http.StatusOK, requestlog.ActionForward) s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) s.get(otherClient, http.StatusOK, requestlog.ActionForward) s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited) s.get(client, http.StatusOK, requestlog.ActionForward) s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned) } // clock is the time a test sets, by which smallwebwaf counts requests and // makes bans. type clock struct { mu sync.Mutex now time.Time } // Now tells the time. func (c *clock) Now() time.Time { c.mu.Lock() defer c.mu.Unlock() return c.now } // advance moves the clock on by d. func (c *clock) advance(d time.Duration) { c.mu.Lock() defer c.mu.Unlock() c.now = c.now.Add(d) } // startWithClock starts smallwebwaf in front of an app that answers 200, // with the settings in env on top of trusting localhost's // X-Forwarded-For, clients' countries looked up at geojsURL, and a clock // set to midnight, the start of a bucket in every window. func startWithClock( t *testing.T, geojsURL string, env map[string]string, ) (*sender, *clock, *proxy.Server) { t.Helper() app := startApp(t, func(http.ResponseWriter, *http.Request) {}) clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)} settings := map[string]string{trustedProxies: trustLocalhost} maps.Copy(settings, env) addr, out, server := startProxyWithClock(t, app.URL, geojsURL, clk.Now, settings) return &sender{t: t, addr: addr, out: out}, clk, server } // sender sends requests to smallwebwaf one after another, each on a // connection of its own, and checks each one's answer and log line. They // must be the only requests smallwebwaf is sent, since the log lines are // matched to them in order. type sender struct { t *testing.T addr string out *output sent int } // get sends a GET request for / from the client at from. func (s *sender) get(from string, status int, action string) logLine { s.t.Helper() return s.request(from, "/", status, action) } // request sends a GET request for path from the client at from, as // X-Forwarded-For names it, and checks that its answer and its log line // have status, 0 for the connection closed without an answer, and that // the line has action. It returns the log line. func (s *sender) request(from, path string, status int, action string) logLine { s.t.Helper() line, _ := s.requestWithHeader(from, path, "", status, action) return line } // requestWithHeader is request with header, such as "Authorization: // Bearer x", added to the request unless it is "". It returns the body of // the answer too. func (s *sender) requestWithHeader( from, path, header string, status int, action string, ) (logLine, string) { s.t.Helper() if header != "" { header += "\r\n" } conn := dial(s.t, s.addr) send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+ "\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+ header+"\r\n") err := conn.SetReadDeadline(time.Now().Add(waitLimit)) if err != nil { s.t.Fatalf("set read deadline: %v", err) } var got answer res, err := http.ReadResponse(bufio.NewReader(conn), nil) switch { case err == nil: got = readAnswer(res) case !errors.Is(err, io.ErrUnexpectedEOF): s.t.Fatalf("read response: %v", err) } _ = conn.Close() if got.status != status { s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, got.status, status) } line := s.out.requestLines(s.t, s.sent+1)[s.sent] s.sent++ wantLine(s.t, line, status, action) return line, string(got.body) }