check / check (push) Waiting to run
Behind a trusted proxy every log line named only the proxy, so abuse could not be traced from webhooker's own logs although the rate limiters already knew the client. The access log, the rate-limit rejection lines, the CSRF warning and the receiver's request line now carry clientIP next to remoteIP. remoteIP still means the connecting peer; clientIP is the address the rate limiters key on, the forwarded client when the peer is inside TRUSTED_PROXIES, worked out once per request by the same code. The README says the field is only as trustworthy as TRUSTED_PROXIES. The access log's 2,560-byte line ceiling holds with the field charged, and a size case with an oversized X-Forwarded-For pins it. 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"])
|
|
}
|
|
})
|
|
}
|
|
}
|
|
}
|