check / check (push) Waiting to run
The access log, the rate-limit rejection lines, the CSRF warning and the receiver's "webhook request received" line now carry clientIP, the address the rate limiters key on, next to remoteIP, the connecting peer. Logging works it out once per request from the same code the rate limiters use and stores it on the request context for the other lines. The CSRF and receiver lines name the peer as remoteIP instead of remote_addr. The README documents the field and that it is only as trustworthy as TRUSTED_PROXIES. Model: opus-5-5
201 lines
5.3 KiB
Go
201 lines
5.3 KiB
Go
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"])
|
|
}
|
|
})
|
|
}
|
|
}
|
|
}
|