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, }, { "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) } }) } }