check / check (push) Successful in 1m9s
realIP took the first X-Forwarded-For entry, which the client itself can write, so behind a proxy that appends to the header a client chose the address dnswatcher logs and the /metrics rate limit counts. It now walks the entries from the right past trusted proxies, using the existing trusted-proxy check, and takes the first that is not one; the leftmost when all are. All X-Forwarded-For header lines are read as one list, since a proxy may add its own line instead of appending to the client's. An empty entry where the client address belongs falls back to the peer address, as an empty first entry did before. X-Real-IP is unchanged. Model: opus-5-5
566 lines
12 KiB
Go
566 lines
12 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 and realIP tests: a client connecting
|
|
// directly, a trusted proxy, and a client behind that proxy as the
|
|
// proxy's X-Real-IP or X-Forwarded-For 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,
|
|
},
|
|
{
|
|
"another client behind the proxy, IPv6-mapped",
|
|
trustedProxy, "::ffff:203.0.113.1",
|
|
trustedProxy, "::ffff: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)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestRealIP checks which address realIP takes as the client's. Each
|
|
// element of forwardedFor is sent as an X-Forwarded-For header line of
|
|
// its own, and 198.51.100.9 is always an entry the client wrote itself.
|
|
func TestRealIP(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
tests := []struct {
|
|
name string
|
|
remoteAddr string
|
|
xRealIP string
|
|
forwardedFor []string
|
|
want string
|
|
}{
|
|
{
|
|
"untrusted peer, both headers ignored",
|
|
directClient, proxiedClient, []string{"198.51.100.9"},
|
|
"198.51.100.1",
|
|
},
|
|
{
|
|
"X-Real-IP from a trusted proxy wins",
|
|
trustedProxy, proxiedClient, []string{"203.0.113.8"},
|
|
proxiedClient,
|
|
},
|
|
{
|
|
"client's own entry, then the one the proxy added",
|
|
trustedProxy, "", []string{"198.51.100.9, 203.0.113.1"},
|
|
proxiedClient,
|
|
},
|
|
{
|
|
"several trusted proxies",
|
|
trustedProxy, "",
|
|
[]string{"198.51.100.9, 203.0.113.1, 10.0.0.3, 10.0.0.2"},
|
|
proxiedClient,
|
|
},
|
|
{
|
|
"proxy adds a header line of its own",
|
|
trustedProxy, "", []string{"198.51.100.9", proxiedClient},
|
|
proxiedClient,
|
|
},
|
|
{
|
|
"every entry a trusted proxy",
|
|
trustedProxy, "", []string{"10.0.0.3, 10.0.0.2"},
|
|
"10.0.0.3",
|
|
},
|
|
{
|
|
"empty where the client address belongs",
|
|
trustedProxy, "", []string{"203.0.113.1, , 10.0.0.2"},
|
|
"10.0.0.1",
|
|
},
|
|
{
|
|
"no headers from a trusted proxy",
|
|
trustedProxy, "", nil,
|
|
"10.0.0.1",
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
req := httptest.NewRequestWithContext(
|
|
t.Context(), http.MethodGet, "/", nil,
|
|
)
|
|
req.RemoteAddr = tt.remoteAddr
|
|
|
|
if tt.xRealIP != "" {
|
|
req.Header.Set("X-Real-IP", tt.xRealIP)
|
|
}
|
|
|
|
for _, line := range tt.forwardedFor {
|
|
req.Header.Add("X-Forwarded-For", line)
|
|
}
|
|
|
|
got := middleware.RealIP(req)
|
|
if got != tt.want {
|
|
t.Errorf("realIP = %q, want %q", got, tt.want)
|
|
}
|
|
})
|
|
}
|
|
}
|