package middleware_test import ( "net/http" "net/http/httptest" "strings" "testing" "github.com/go-chi/chi/v5" "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(nil, notify.Params{ Logger: log, Config: &config.Config{}, }) if err != nil { t.Fatalf("notify.New: %v", err) } hnd, err := handlers.New(nil, handlers.Params{ Logger: log, Globals: glob, State: state.NewForTest(), 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) } }