check / check (push) Successful in 1m9s
/metrics is behind a password, and REPO_POLICIES.md requires rate limiting on password logins. Each client address may now send it 30 requests a minute, counted by httprate before Basic Auth, so failed logins use up the allowance and a request over it gets 429 without the password being checked. The address is the one the existing trusted-proxy logic in internal/middleware works out, with IPv6 addresses grouped by /64. A Prometheus server scraping every 15 seconds sends 4 requests a minute. Model: opus-5-5
479 lines
10 KiB
Go
479 lines
10 KiB
Go
package middleware_test
|
|
|
|
import (
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/go-chi/chi/v5"
|
|
"go.uber.org/fx/fxtest"
|
|
|
|
"sneak.berlin/go/dnswatcher/internal/config"
|
|
"sneak.berlin/go/dnswatcher/internal/globals"
|
|
"sneak.berlin/go/dnswatcher/internal/handlers"
|
|
"sneak.berlin/go/dnswatcher/internal/logger"
|
|
"sneak.berlin/go/dnswatcher/internal/middleware"
|
|
"sneak.berlin/go/dnswatcher/internal/notify"
|
|
"sneak.berlin/go/dnswatcher/internal/state"
|
|
)
|
|
|
|
// Expected security header values, spelled out literally so that any
|
|
// change to the middleware has to be made deliberately here as well.
|
|
const (
|
|
wantHSTS = "max-age=31536000; includeSubDomains"
|
|
|
|
wantCSP = "default-src 'self'; " +
|
|
"script-src 'none'; " +
|
|
"style-src 'self'; " +
|
|
"img-src 'self'; " +
|
|
"font-src 'none'; " +
|
|
"connect-src 'none'; " +
|
|
"object-src 'none'; " +
|
|
"base-uri 'none'; " +
|
|
"form-action 'none'; " +
|
|
"frame-ancestors 'none'"
|
|
|
|
wantFrameOptions = "DENY"
|
|
|
|
wantContentTypeOptions = "nosniff"
|
|
|
|
wantReferrerPolicy = "no-referrer"
|
|
|
|
wantPermissionsPolicy = "accelerometer=(), " +
|
|
"autoplay=(), " +
|
|
"camera=(), " +
|
|
"display-capture=(), " +
|
|
"encrypted-media=(), " +
|
|
"fullscreen=(), " +
|
|
"geolocation=(), " +
|
|
"gyroscope=(), " +
|
|
"magnetometer=(), " +
|
|
"microphone=(), " +
|
|
"midi=(), " +
|
|
"payment=(), " +
|
|
"picture-in-picture=(), " +
|
|
"publickey-credentials-get=(), " +
|
|
"screen-wake-lock=(), " +
|
|
"usb=(), " +
|
|
"xr-spatial-tracking=()"
|
|
)
|
|
|
|
// stylesheetPath is the only subresource the dashboard loads.
|
|
const stylesheetPath = "/s/css/tailwind.min.css"
|
|
|
|
// newTestLogger builds a logger for direct component construction.
|
|
func newTestLogger(t *testing.T) *logger.Logger {
|
|
t.Helper()
|
|
|
|
glob, err := globals.New(nil)
|
|
if err != nil {
|
|
t.Fatalf("globals.New: %v", err)
|
|
}
|
|
|
|
log, err := logger.New(nil, logger.Params{Globals: glob})
|
|
if err != nil {
|
|
t.Fatalf("logger.New: %v", err)
|
|
}
|
|
|
|
return log
|
|
}
|
|
|
|
// newTestMiddleware builds a Middleware without an fx application.
|
|
func newTestMiddleware(t *testing.T) *middleware.Middleware {
|
|
t.Helper()
|
|
|
|
glob, err := globals.New(nil)
|
|
if err != nil {
|
|
t.Fatalf("globals.New: %v", err)
|
|
}
|
|
|
|
mw, err := middleware.New(nil, middleware.Params{
|
|
Logger: newTestLogger(t),
|
|
Globals: glob,
|
|
Config: &config.Config{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("middleware.New: %v", err)
|
|
}
|
|
|
|
return mw
|
|
}
|
|
|
|
// serveWithSecurityHeaders runs a GET through SecurityHeaders and
|
|
// returns the recorded response.
|
|
func serveWithSecurityHeaders(
|
|
t *testing.T,
|
|
target string,
|
|
handler http.Handler,
|
|
) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
|
|
mw := newTestMiddleware(t)
|
|
rec := httptest.NewRecorder()
|
|
req := httptest.NewRequestWithContext(
|
|
t.Context(), http.MethodGet, target, nil,
|
|
)
|
|
|
|
mw.SecurityHeaders()(handler).ServeHTTP(rec, req)
|
|
|
|
return rec
|
|
}
|
|
|
|
// okHandler writes a trivial 200 response.
|
|
func okHandler() http.Handler {
|
|
return http.HandlerFunc(func(
|
|
writer http.ResponseWriter,
|
|
_ *http.Request,
|
|
) {
|
|
writer.WriteHeader(http.StatusOK)
|
|
})
|
|
}
|
|
|
|
func TestSecurityHeaders(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
tests := []struct {
|
|
name string
|
|
header string
|
|
want string
|
|
}{
|
|
{
|
|
"hsts",
|
|
"Strict-Transport-Security",
|
|
wantHSTS,
|
|
},
|
|
{
|
|
"csp",
|
|
"Content-Security-Policy",
|
|
wantCSP,
|
|
},
|
|
{
|
|
"frame options",
|
|
"X-Frame-Options",
|
|
wantFrameOptions,
|
|
},
|
|
{
|
|
"content type options",
|
|
"X-Content-Type-Options",
|
|
wantContentTypeOptions,
|
|
},
|
|
{
|
|
"referrer policy",
|
|
"Referrer-Policy",
|
|
wantReferrerPolicy,
|
|
},
|
|
{
|
|
"permissions policy",
|
|
"Permissions-Policy",
|
|
wantPermissionsPolicy,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
rec := serveWithSecurityHeaders(t, "/", okHandler())
|
|
|
|
got := rec.Header().Get(tt.header)
|
|
if got != tt.want {
|
|
t.Errorf(
|
|
"%s = %q, want %q",
|
|
tt.header, got, tt.want,
|
|
)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestSecurityHeadersCSPDirectives guards the properties the repo
|
|
// policy requires of the content security policy itself.
|
|
func TestSecurityHeadersCSPDirectives(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
rec := serveWithSecurityHeaders(t, "/", okHandler())
|
|
csp := rec.Header().Get("Content-Security-Policy")
|
|
|
|
forbidden := []string{"unsafe-inline", "unsafe-eval"}
|
|
for _, directive := range forbidden {
|
|
if strings.Contains(csp, directive) {
|
|
t.Errorf("CSP must not contain %q: %q", directive, csp)
|
|
}
|
|
}
|
|
|
|
required := []string{
|
|
"default-src 'self'",
|
|
"script-src 'none'",
|
|
"style-src 'self'",
|
|
"frame-ancestors 'none'",
|
|
}
|
|
for _, directive := range required {
|
|
if !strings.Contains(csp, directive) {
|
|
t.Errorf("CSP must contain %q: %q", directive, csp)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestSecurityHeadersOnErrorResponse verifies the headers are emitted
|
|
// even when the wrapped handler fails, since they are set before the
|
|
// handler runs.
|
|
func TestSecurityHeadersOnErrorResponse(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
failing := http.HandlerFunc(func(
|
|
writer http.ResponseWriter,
|
|
_ *http.Request,
|
|
) {
|
|
http.Error(
|
|
writer,
|
|
"boom",
|
|
http.StatusInternalServerError,
|
|
)
|
|
})
|
|
|
|
rec := serveWithSecurityHeaders(t, "/api/v1/status", failing)
|
|
|
|
if rec.Code != http.StatusInternalServerError {
|
|
t.Fatalf("status = %d, want 500", rec.Code)
|
|
}
|
|
|
|
if got := rec.Header().Get(
|
|
"X-Content-Type-Options",
|
|
); got != wantContentTypeOptions {
|
|
t.Errorf(
|
|
"X-Content-Type-Options = %q, want %q",
|
|
got, wantContentTypeOptions,
|
|
)
|
|
}
|
|
|
|
if got := rec.Header().Get(
|
|
"Strict-Transport-Security",
|
|
); got != wantHSTS {
|
|
t.Errorf(
|
|
"Strict-Transport-Security = %q, want %q",
|
|
got, wantHSTS,
|
|
)
|
|
}
|
|
}
|
|
|
|
// newTestHandlers builds real Handlers with empty monitoring state.
|
|
func newTestHandlers(t *testing.T) *handlers.Handlers {
|
|
t.Helper()
|
|
|
|
glob, err := globals.New(nil)
|
|
if err != nil {
|
|
t.Fatalf("globals.New: %v", err)
|
|
}
|
|
|
|
log := newTestLogger(t)
|
|
|
|
notifier, err := notify.New(fxtest.NewLifecycle(t), notify.Params{
|
|
Logger: log,
|
|
Config: &config.Config{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("notify.New: %v", err)
|
|
}
|
|
|
|
st, err := state.New(fxtest.NewLifecycle(t), state.Params{
|
|
Logger: log,
|
|
Config: &config.Config{DataDir: t.TempDir()},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("state.New: %v", err)
|
|
}
|
|
|
|
hnd, err := handlers.New(nil, handlers.Params{
|
|
Logger: log,
|
|
Globals: glob,
|
|
State: st,
|
|
Notify: notifier,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("handlers.New: %v", err)
|
|
}
|
|
|
|
return hnd
|
|
}
|
|
|
|
// TestDashboardRendersWithSecurityHeaders renders the real dashboard
|
|
// through the middleware and checks that the policy still permits the
|
|
// one stylesheet the page loads.
|
|
func TestDashboardRendersWithSecurityHeaders(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
mw := newTestMiddleware(t)
|
|
hnd := newTestHandlers(t)
|
|
|
|
router := chi.NewRouter()
|
|
router.Use(mw.SecurityHeaders())
|
|
router.Get("/", hnd.HandleDashboard())
|
|
|
|
rec := httptest.NewRecorder()
|
|
req := httptest.NewRequestWithContext(
|
|
t.Context(), http.MethodGet, "/", nil,
|
|
)
|
|
|
|
router.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusOK {
|
|
t.Fatalf("status = %d, want 200", rec.Code)
|
|
}
|
|
|
|
body := rec.Body.String()
|
|
if !strings.Contains(body, stylesheetPath) {
|
|
t.Errorf("dashboard does not reference %q", stylesheetPath)
|
|
}
|
|
|
|
if !strings.Contains(body, "dnswatcher") {
|
|
t.Errorf("dashboard body looks empty: %d bytes", len(body))
|
|
}
|
|
|
|
csp := rec.Header().Get("Content-Security-Policy")
|
|
if csp != wantCSP {
|
|
t.Errorf("CSP = %q, want %q", csp, wantCSP)
|
|
}
|
|
|
|
// The stylesheet is same-origin, so style-src 'self' allows it.
|
|
if !strings.Contains(csp, "style-src 'self'") {
|
|
t.Errorf("CSP would block %q: %q", stylesheetPath, csp)
|
|
}
|
|
}
|
|
|
|
// Addresses for the rate limit tests: a client connecting directly, a
|
|
// trusted proxy, and a client behind that proxy as its X-Real-IP
|
|
// header names it.
|
|
const (
|
|
directClient = "198.51.100.1:4000"
|
|
trustedProxy = "10.0.0.1:4000"
|
|
proxiedClient = "203.0.113.1"
|
|
)
|
|
|
|
// statusFrom sends a GET through handler as if from remoteAddr, with
|
|
// an X-Real-IP header when xRealIP is not empty, and returns the
|
|
// response status.
|
|
func statusFrom(
|
|
t *testing.T,
|
|
handler http.Handler,
|
|
remoteAddr string,
|
|
xRealIP string,
|
|
) int {
|
|
t.Helper()
|
|
|
|
req := httptest.NewRequestWithContext(
|
|
t.Context(), http.MethodGet, "/metrics", nil,
|
|
)
|
|
req.RemoteAddr = remoteAddr
|
|
|
|
if xRealIP != "" {
|
|
req.Header.Set("X-Real-IP", xRealIP)
|
|
}
|
|
|
|
rec := httptest.NewRecorder()
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
return rec.Code
|
|
}
|
|
|
|
// TestMetricsRateLimitAllowsScraping checks that one address can send,
|
|
// within one window, what two Prometheus servers scraping every 5
|
|
// seconds send in that time, without being turned away.
|
|
func TestMetricsRateLimitAllowsScraping(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
const scrapeInterval = 5 * time.Second
|
|
|
|
scrapes := 2 * int(middleware.MetricsRequestWindow/scrapeInterval)
|
|
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
|
|
|
|
for i := range scrapes {
|
|
got := statusFrom(t, limited, directClient, "")
|
|
if got != http.StatusOK {
|
|
t.Fatalf(
|
|
"scrape %d of %d: status = %d, want 200",
|
|
i+1, scrapes, got,
|
|
)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestMetricsRateLimitKeysOnClientAddress checks which requests share
|
|
// an allowance. Each case uses up the allowance of one client, then
|
|
// sends one more request.
|
|
func TestMetricsRateLimitKeysOnClientAddress(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
tests := []struct {
|
|
name string
|
|
usedRemoteAddr string
|
|
usedXRealIP string
|
|
nextRemoteAddr string
|
|
nextXRealIP string
|
|
want int
|
|
}{
|
|
{
|
|
"same address",
|
|
directClient, "",
|
|
directClient, "",
|
|
http.StatusTooManyRequests,
|
|
},
|
|
{
|
|
"another address",
|
|
directClient, "",
|
|
"198.51.100.2:4000", "",
|
|
http.StatusOK,
|
|
},
|
|
{
|
|
"own X-Real-IP from an untrusted address",
|
|
directClient, "",
|
|
directClient, "203.0.113.9",
|
|
http.StatusTooManyRequests,
|
|
},
|
|
{
|
|
"same client behind the proxy",
|
|
trustedProxy, proxiedClient,
|
|
trustedProxy, proxiedClient,
|
|
http.StatusTooManyRequests,
|
|
},
|
|
{
|
|
"another client behind the proxy",
|
|
trustedProxy, proxiedClient,
|
|
trustedProxy, "203.0.113.2",
|
|
http.StatusOK,
|
|
},
|
|
{
|
|
"same IPv6 /64",
|
|
"[2001:db8::1]:4000", "",
|
|
"[2001:db8::2]:4000", "",
|
|
http.StatusTooManyRequests,
|
|
},
|
|
{
|
|
"another IPv6 /64",
|
|
"[2001:db8::1]:4000", "",
|
|
"[2001:db8:0:1::1]:4000", "",
|
|
http.StatusOK,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
|
|
|
|
for range middleware.MetricsRequestLimit {
|
|
statusFrom(t, limited, tt.usedRemoteAddr, tt.usedXRealIP)
|
|
}
|
|
|
|
got := statusFrom(
|
|
t, limited, tt.nextRemoteAddr, tt.nextXRealIP,
|
|
)
|
|
if got != tt.want {
|
|
t.Errorf("status = %d, want %d", got, tt.want)
|
|
}
|
|
})
|
|
}
|
|
}
|