package middleware_test import ( "context" "fmt" "log/slog" "net/http" "net/http/httptest" "net/netip" "os" "testing" "github.com/stretchr/testify/assert" "sneak.berlin/go/webhooker/internal/config" "sneak.berlin/go/webhooker/internal/middleware" ) func TestLoginRateLimit_AllowsGET(t *testing.T) { t.Parallel() m, _ := testMiddleware(t, config.EnvironmentDev) var callCount int handler := m.LoginRateLimit()(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { callCount++ w.WriteHeader(http.StatusOK) }, )) // GET requests should never be rate-limited for i := range 20 { req := httptest.NewRequestWithContext( context.Background(), http.MethodGet, "/pages/login", nil, ) req.RemoteAddr = "192.168.1.1:12345" w := httptest.NewRecorder() handler.ServeHTTP(w, req) assert.Equal( t, http.StatusOK, w.Code, "GET request %d should pass", i, ) } assert.Equal(t, 20, callCount) } // runPostLimitTest exercises a POST-only rate limit middleware: // the first limit POSTs to path from ip must pass, and the next // one must be rejected with 429 without reaching the handler. func runPostLimitTest( t *testing.T, mw func(http.Handler) http.Handler, limit int, path, ip string, ) { t.Helper() var callCount int handler := mw(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { callCount++ w.WriteHeader(http.StatusOK) }, )) // The first limit POST requests should succeed for i := range limit { req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, path, nil, ) req.RemoteAddr = ip w := httptest.NewRecorder() handler.ServeHTTP(w, req) assert.Equal( t, http.StatusOK, w.Code, "POST request %d should pass", i, ) } // Next POST should be rate-limited req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, path, nil, ) req.RemoteAddr = ip w := httptest.NewRecorder() handler.ServeHTTP(w, req) assert.Equal( t, http.StatusTooManyRequests, w.Code, "POST after limit should be 429", ) assert.Equal(t, limit, callCount) } func TestLoginRateLimit_LimitsPOST(t *testing.T) { t.Parallel() m, _ := testMiddleware(t, config.EnvironmentDev) runPostLimitTest( t, m.LoginRateLimit(), middleware.LoginRateLimitConst, "/pages/login", "10.0.0.1:12345", ) } func TestPasswordChangeRateLimit_LimitsPOST(t *testing.T) { t.Parallel() m, _ := testMiddleware(t, config.EnvironmentDev) runPostLimitTest( t, m.PasswordChangeRateLimit(), middleware.PasswordChangeRateLimitConst, "/user/admin/password", "10.0.0.2:12345", ) } func TestLoginRateLimit_IndependentPerIP(t *testing.T) { t.Parallel() m, _ := testMiddleware(t, config.EnvironmentDev) handler := m.LoginRateLimit()(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }, )) // Exhaust limit for IP1 for range middleware.LoginRateLimitConst { req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, "/pages/login", nil, ) req.RemoteAddr = "1.2.3.4:12345" w := httptest.NewRecorder() handler.ServeHTTP(w, req) } // IP1 should be rate-limited req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, "/pages/login", nil, ) req.RemoteAddr = "1.2.3.4:12345" w := httptest.NewRecorder() handler.ServeHTTP(w, req) assert.Equal(t, http.StatusTooManyRequests, w.Code) // IP2 should still be allowed req2 := httptest.NewRequestWithContext( context.Background(), http.MethodPost, "/pages/login", nil, ) req2.RemoteAddr = "5.6.7.8:12345" w2 := httptest.NewRecorder() handler.ServeHTTP(w2, req2) assert.Equal( t, http.StatusOK, w2.Code, "different IP should not be affected", ) } // okHandler is the terminal handler the limiter middleware wraps // in these tests: it answers 200 to anything that reaches it. func okHandler() http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }, ) } // rateLimitMiddleware builds a Middleware around cfg, whose // TrustedProxies field is what the rate limit key function gates // forwarded-header trust on. func rateLimitMiddleware( t *testing.T, cfg *config.Config, ) *middleware.Middleware { t.Helper() log := slog.New(slog.NewTextHandler( os.Stderr, &slog.HandlerOptions{Level: slog.LevelDebug}, )) return middleware.NewForTest(log, cfg, nil) } // trustedProxies parses CIDR strings for a test Config. func trustedProxies(cidrs ...string) []netip.Prefix { prefixes := make([]netip.Prefix, 0, len(cidrs)) for _, cidr := range cidrs { prefixes = append(prefixes, netip.MustParsePrefix(cidr)) } return prefixes } // postWithHeaders sends one POST to the handler from peer with the // given headers set and returns the recorder. func postWithHeaders( handler http.Handler, peer, path string, headers map[string]string, ) *httptest.ResponseRecorder { req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, path, nil, ) req.RemoteAddr = peer for name, value := range headers { req.Header.Set(name, value) } w := httptest.NewRecorder() handler.ServeHTTP(w, req) return w } // receiverLimitedHandler builds a ReceiverRateLimit-wrapped // handler with the given per-minute limit and no trusted proxies. func receiverLimitedHandler( t *testing.T, limit int, ) http.Handler { t.Helper() m := rateLimitMiddleware( t, &config.Config{ReceiverRateLimit: limit}, ) return m.ReceiverRateLimit()(okHandler()) } // receiverPost sends one POST to the handler from the given IP // and path and returns the recorder. func receiverPost( handler http.Handler, ip, path string, ) *httptest.ResponseRecorder { req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, path, nil, ) req.RemoteAddr = ip w := httptest.NewRecorder() handler.ServeHTTP(w, req) return w } func TestReceiverRateLimit_LimitsPerIPAndPath(t *testing.T) { t.Parallel() const limit = 3 handler := receiverLimitedHandler(t, limit) // The first limit requests from one IP to one entrypoint // pass. for i := range limit { w := receiverPost( handler, "9.9.9.9:1234", "/webhook/uuid-a", ) assert.Equal( t, http.StatusOK, w.Code, "request %d should pass", i, ) } // The next request over the limit is rejected with a 429 // carrying a Retry-After header. w := receiverPost( handler, "9.9.9.9:1234", "/webhook/uuid-a", ) assert.Equal(t, http.StatusTooManyRequests, w.Code) assert.NotEmpty( t, w.Header().Get("Retry-After"), "429 must carry a Retry-After header", ) // The same IP is not limited on a different entrypoint. w = receiverPost( handler, "9.9.9.9:1234", "/webhook/uuid-b", ) assert.Equal( t, http.StatusOK, w.Code, "a different entrypoint must not be affected", ) // A different IP is not limited on the same entrypoint. w = receiverPost( handler, "8.8.8.8:1234", "/webhook/uuid-a", ) assert.Equal( t, http.StatusOK, w.Code, "a different client IP must not be affected", ) } // TestReceiverRateLimit_CountsEveryMethod proves the receiver // limit counts non-POST requests too: a GET shares the bucket // with a POST and is itself rejected once over the limit. func TestReceiverRateLimit_CountsEveryMethod(t *testing.T) { t.Parallel() const ( limit = 2 ip = "7.7.7.7:1234" path = "/webhook/uuid-c" ) handler := receiverLimitedHandler(t, limit) get := func() *httptest.ResponseRecorder { req := httptest.NewRequestWithContext( context.Background(), http.MethodGet, path, nil, ) req.RemoteAddr = ip w := httptest.NewRecorder() handler.ServeHTTP(w, req) return w } // One POST plus one GET fill the bucket, so the GET must // have been counted. assert.Equal( t, http.StatusOK, receiverPost(handler, ip, path).Code, ) assert.Equal(t, http.StatusOK, get().Code) assert.Equal( t, http.StatusTooManyRequests, get().Code, "a GET over the limit must be rate-limited", ) } const ( loginPath = "/pages/login" headerXFF = "X-Forwarded-For" headerReal = "X-Real-IP" headerTrue = "True-Client-IP" ) // TestRateLimitKey_SpoofedForwardedFromUntrustedPeer is the test // this gating exists for: with no trusted proxies configured (the // default), a client that rotates a forwarded header on every // request must stay in one bucket. If forwarded headers were // trusted unconditionally, each spoofed value would mint a fresh // bucket and the limit would stop no one. func TestRateLimitKey_SpoofedForwardedFromUntrustedPeer( t *testing.T, ) { t.Parallel() for _, header := range []string{ headerXFF, headerReal, headerTrue, } { t.Run(header, func(t *testing.T) { t.Parallel() m := rateLimitMiddleware(t, &config.Config{}) handler := m.LoginRateLimit()(okHandler()) const peer = "203.0.113.9:44444" for i := range middleware.LoginRateLimitConst { w := postWithHeaders( handler, peer, loginPath, map[string]string{ header: fmt.Sprintf( "198.51.100.%d", i+1, ), }, ) assert.Equal( t, http.StatusOK, w.Code, "request %d should pass", i, ) } w := postWithHeaders( handler, peer, loginPath, map[string]string{header: "198.51.100.200"}, ) assert.Equal( t, http.StatusTooManyRequests, w.Code, "a spoofed %s from an untrusted peer must "+ "not mint a fresh bucket", header, ) }) } } // TestRateLimitKey_ForwardedHonouredFromTrustedPeer checks the // other half: when the direct peer is a configured trusted proxy, // the forwarded client address is what buckets are keyed on, so // one sender behind the proxy cannot exhaust another's limit. func TestRateLimitKey_ForwardedHonouredFromTrustedPeer( t *testing.T, ) { t.Parallel() m := rateLimitMiddleware(t, &config.Config{ TrustedProxies: trustedProxies("10.0.0.0/8"), }) handler := m.LoginRateLimit()(okHandler()) const peer = "10.0.0.1:44444" first := map[string]string{headerXFF: "198.51.100.7"} for range middleware.LoginRateLimitConst { postWithHeaders(handler, peer, loginPath, first) } w := postWithHeaders(handler, peer, loginPath, first) assert.Equal( t, http.StatusTooManyRequests, w.Code, "the forwarded client's own bucket must fill up", ) w = postWithHeaders( handler, peer, loginPath, map[string]string{headerXFF: "198.51.100.8"}, ) assert.Equal( t, http.StatusOK, w.Code, "a forwarded header from a trusted peer must be honoured", ) } // TestRateLimitKey_ChainWalkSkipsClientPrepended covers the // residual spoofing route behind a trusted proxy: the client // controls the leftmost X-Forwarded-For entries, so the key is the // rightmost hop that is not itself trusted. Rotating the prepended // entry must not create new buckets. func TestRateLimitKey_ChainWalkSkipsClientPrepended(t *testing.T) { t.Parallel() m := rateLimitMiddleware(t, &config.Config{ TrustedProxies: trustedProxies("10.0.0.0/8"), }) handler := m.LoginRateLimit()(okHandler()) const peer = "10.0.0.1:44444" chain := func(spoof string) map[string]string { return map[string]string{ headerXFF: spoof + ", 198.51.100.7, 10.0.0.2", } } for i := range middleware.LoginRateLimitConst { w := postWithHeaders( handler, peer, loginPath, chain(fmt.Sprintf("9.9.9.%d", i+1)), ) assert.Equal( t, http.StatusOK, w.Code, "request %d should pass", i, ) } w := postWithHeaders( handler, peer, loginPath, chain("9.9.9.200"), ) assert.Equal( t, http.StatusTooManyRequests, w.Code, "a client-prepended X-Forwarded-For entry must not "+ "mint a fresh bucket", ) } // TestReceiverRateLimit_IgnoresForwardedFromUntrustedPeer proves // the receiver limiter uses the same gated key function as the // POST limiters. func TestReceiverRateLimit_IgnoresForwardedFromUntrustedPeer( t *testing.T, ) { t.Parallel() const ( limit = 3 peer = "203.0.113.10:44444" path = "/webhook/uuid-d" ) handler := receiverLimitedHandler(t, limit) for i := range limit { w := postWithHeaders( handler, peer, path, map[string]string{ headerXFF: fmt.Sprintf( "198.51.100.%d", i+1, ), }, ) assert.Equal( t, http.StatusOK, w.Code, "request %d should pass", i, ) } w := postWithHeaders( handler, peer, path, map[string]string{headerXFF: "198.51.100.200"}, ) assert.Equal( t, http.StatusTooManyRequests, w.Code, "a spoofed X-Forwarded-For from an untrusted peer must "+ "not mint a fresh receiver bucket", ) }