package middleware_test import ( "bytes" "context" "log/slog" "net/http" "net/http/httptest" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "sneak.berlin/go/webhooker/internal/config" "sneak.berlin/go/webhooker/internal/middleware" ) const ( // forwardedChain is the X-Forwarded-For a request arrives with: // the client, then a second proxy inside trustedProxyCIDR that the // request passed through before reaching trustedPeer. forwardedChain = clientIPv4 + ", 10.0.0.2" // untrustedPeer is a peer outside trustedProxyCIDR, so its // X-Forwarded-For is ignored and the peer is the client. untrustedPeer = "192.0.2.10:5555" // oneRequestPerMinute is the receiver limit these tests install: // the second request on a path is rejected, and the aggregate // limit is ReceiverAggregateMultiplierConst. oneRequestPerMinute = 1 ) // clientLogSite is one log line that names the client. build wraps the // middleware that writes it around a handler, and requests is how many // identical requests it takes before the line is written. type clientLogSite struct { build func(m *middleware.Middleware) http.Handler requests int } // clientLogSites maps the message of each line that names the client // to the way to make it be written. func clientLogSites() map[string]clientLogSite { served := func(*middleware.Middleware) http.Handler { return okHandler() } receiver := func(m *middleware.Middleware) http.Handler { return m.ReceiverRateLimit()(okHandler()) } login := func(m *middleware.Middleware) http.Handler { return http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { m.RecordLoginFailure(r, "someone") w.WriteHeader(http.StatusUnauthorized) }, ) } csrf := func(m *middleware.Middleware) http.Handler { return m.CSRF(http.HandlerFunc(forbidden))(okHandler()) } passwordChange := func(m *middleware.Middleware) http.Handler { return m.PasswordChangeRateLimit()(okHandler()) } replay := func(m *middleware.Middleware) http.Handler { return m.ReplayRateLimit()(okHandler()) } resubmit := func(m *middleware.Middleware) http.Handler { return m.ResubmitRateLimit()(okHandler()) } return map[string]clientLogSite{ "http request": { build: served, requests: 1, }, "webhook receiver rate limit exceeded": { build: receiver, requests: oneRequestPerMinute + 1, }, // The aggregate limit sits in front of the per-entrypoint // one, so the requests that one rejects count towards it. "webhook receiver aggregate rate limit exceeded": { build: receiver, requests: middleware.ReceiverAggregateMultiplierConst* oneRequestPerMinute + 1, }, "login failure limit exceeded": { build: login, requests: middleware.LoginRateLimitConst + 1, }, "csrf: token validation failed": { build: csrf, requests: 1, }, "password change rate limit exceeded": { build: passwordChange, requests: middleware.PasswordChangeRateLimitConst + 1, }, "delivery replay rate limit exceeded": { build: replay, requests: middleware.ReplayRateLimitConst + 1, }, "event resubmit rate limit exceeded": { build: resubmit, requests: middleware.ResubmitRateLimitConst + 1, }, } } // clientLogLines sends the site's requests from peer, each carrying // forwardedChain, through Logging and then the site, as production // does, and returns the logged lines whose message is msg. func clientLogLines( t *testing.T, site clientLogSite, msg, peer string, ) []map[string]any { t.Helper() buf := new(bytes.Buffer) log := slog.New(slog.NewJSONHandler( buf, &slog.HandlerOptions{Level: slog.LevelDebug}, )) cfg := &config.Config{ Environment: config.EnvironmentDev, ReceiverRateLimit: oneRequestPerMinute, TrustedProxies: trustedProxies(trustedProxyCIDR), } m := middleware.NewForTest( log, cfg, newTestSessionManager(cfg, log, nil), ) handler := m.Logging()(site.build(m)) for range site.requests { req := httptest.NewRequestWithContext( context.Background(), http.MethodPost, "/h/x", nil, ) req.RemoteAddr = peer req.Header.Set(headerXFF, forwardedChain) handler.ServeHTTP(httptest.NewRecorder(), req) } var lines []map[string]any for _, entry := range accessLogEntries(t, buf) { if entry["msg"] == msg { lines = append(lines, entry) } } return lines } // TestClientIP_LoggedNextToThePeer checks that every line that names // the client carries both addresses: remoteIP, the connecting peer, // and clientIP, the client the rate limiters key on. func TestClientIP_LoggedNextToThePeer(t *testing.T) { t.Parallel() cases := map[string]struct { peer string wantRemote string wantClient string }{ "trusted proxy with a forwarded chain": { peer: trustedPeer, wantRemote: "10.0.0.1", wantClient: clientIPv4, }, "untrusted peer": { peer: untrustedPeer, wantRemote: "192.0.2.10", wantClient: "192.0.2.10", }, } for msg, site := range clientLogSites() { for name, tc := range cases { t.Run(msg+"/"+name, func(t *testing.T) { t.Parallel() lines := clientLogLines(t, site, msg, tc.peer) require.NotEmpty(t, lines, "%q was never logged", msg) for _, line := range lines { assert.Equal(t, tc.wantRemote, line["remoteIP"]) assert.Equal(t, tc.wantClient, line["clientIP"]) } }) } } }