check / check (push) Waiting to run
The test helpers lived in ordinary `testing.go` files inside the config, database, middleware and session packages, so they were built into the binary and the shared `test-support` lint rule could not see them. The four files are gone: the session's helpers move into its own `_test.go` file, and the rest into `configtest`, `databasetest` and `middlewaretest`, which the `depguard` deny list now names, so a non-test file importing them fails lint. The test-support packages build through the production constructors. Judgement call: the session, the middleware and the webhook database manager now take the plain logger they log through, which the application wiring provides. Judgement call: two idle-expiry tests move the stored timestamps back instead of advancing a fake clock. Model: opus-5-5
202 lines
5.4 KiB
Go
202 lines
5.4 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"
|
|
"sneak.berlin/go/webhooker/internal/middleware/middlewaretest"
|
|
)
|
|
|
|
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 := middlewaretest.New(
|
|
t, log, cfg, newTestSessionManager(t, cfg),
|
|
)
|
|
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"])
|
|
}
|
|
})
|
|
}
|
|
}
|
|
}
|