check / check (push) Waiting to run
The four testing.go files in config, database, middleware and session are gone. ClearEnvForTest moves to internal/config/configtest as ClearEnv; the webhook database manager helpers move to internal/database/databasetest as NewWebhookDBManager and NewWebhookDBManagerWithLogger, and the middleware's NewForTest to internal/middleware/middlewaretest as New, both now built through the production constructors. Tests that wrapped an open main database use database.Open. The session helpers move into the session package's export_test.go; the middleware tests build their session through session.New and age its timestamps instead of using a fake clock. The session, middleware and webhook database manager now take the *slog.Logger they log through, so tests in other packages can give them their own. 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"])
|
|
}
|
|
})
|
|
}
|
|
}
|
|
}
|