package middleware import ( "bytes" "log/slog" "net/http" "net/http/httptest" "net/url" "strings" "testing" "github.com/prometheus/client_golang/prometheus/promhttp" "sneak.berlin/go/pixa/internal/config" ) // TestCORSAnswersWithConfiguredOrigin checks that the CORS middleware // answers with access_control_allow_origin, where "*" lets any origin read // responses and a single origin lets only that origin read them. func TestCORSAnswersWithConfiguredOrigin(t *testing.T) { t.Parallel() const appOrigin = "https://app.example.com" cases := []struct { configured string requestOrigin string want string }{ {"*", "https://any.example.com", "*"}, {appOrigin, appOrigin, appOrigin}, {appOrigin, "https://other.example.com", ""}, } testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) for _, tc := range cases { mw := &Middleware{ log: slog.Default(), config: &config.Config{AccessControlAllowOrigin: tc.configured}, } handler := mw.CORS()(testHandler) req := httptest.NewRequestWithContext( t.Context(), http.MethodGet, "/v1/image/example.com/a.jpg/1x1.png", nil) req.Header.Set("Origin", tc.requestOrigin) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) got := rec.Header().Get("Access-Control-Allow-Origin") if got != tc.want { t.Errorf("configured %q, request from %q: "+ "Access-Control-Allow-Origin = %q, want %q", tc.configured, tc.requestOrigin, got, tc.want) } } } // TestCORSAnswersPreflightWithConfiguredOrigin checks that the CORS // middleware answers a preflight request, which the CORS library handles // apart from other requests, the same way: "*" lets any origin read // responses and a single origin lets only that origin read them. func TestCORSAnswersPreflightWithConfiguredOrigin(t *testing.T) { t.Parallel() const appOrigin = "https://app.example.com" cases := []struct { configured string requestOrigin string want string }{ {"*", "https://any.example.com", "*"}, {appOrigin, appOrigin, appOrigin}, {appOrigin, "https://other.example.com", ""}, } testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) for _, tc := range cases { mw := &Middleware{ log: slog.Default(), config: &config.Config{AccessControlAllowOrigin: tc.configured}, } handler := mw.CORS()(testHandler) // An OPTIONS request naming the method it asks about is the // preflight a browser sends before some cross-origin requests. req := httptest.NewRequestWithContext( t.Context(), http.MethodOptions, "/v1/image/example.com/a.jpg/1x1.png", nil) req.Header.Set("Origin", tc.requestOrigin) req.Header.Set("Access-Control-Request-Method", http.MethodGet) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) got := rec.Header().Get("Access-Control-Allow-Origin") if got != tc.want { t.Errorf("configured %q, preflight from %q: "+ "Access-Control-Allow-Origin = %q, want %q", tc.configured, tc.requestOrigin, got, tc.want) } } } // TestMetricsAuthRequiresConfiguredCredentials checks that MetricsAuth on // its own answers 401 with a challenge to a request without credentials or // with a wrong username or password, and lets a request with the configured // username and password through. That the router puts it in front of // /metrics is not tested. func TestMetricsAuthRequiresConfiguredCredentials(t *testing.T) { t.Parallel() const ( username = "metricsuser" password = "metricspass" challenge = `Basic realm="metrics"` ) // An empty username stands for a request sent without credentials. cases := []struct { name string username string password string wantReached bool }{ {"no credentials", "", "", false}, {"wrong username", "someone", password, false}, {"wrong password", username, "wrongpass", false}, {"configured credentials", username, password, true}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { t.Parallel() mw := &Middleware{ log: slog.Default(), config: &config.Config{ MetricsUsername: username, MetricsPassword: password, }, } reached := false handler := mw.MetricsAuth()(http.HandlerFunc( func(http.ResponseWriter, *http.Request) { reached = true })) req := httptest.NewRequestWithContext( t.Context(), http.MethodGet, "/metrics", nil) if tc.username != "" { req.SetBasicAuth(tc.username, tc.password) } rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) if reached != tc.wantReached { t.Fatalf("request reached /metrics = %v, want %v", reached, tc.wantReached) } if tc.wantReached { return } if rec.Code != http.StatusUnauthorized { t.Errorf("status = %d, want %d", rec.Code, http.StatusUnauthorized) } if got := rec.Header().Get("WWW-Authenticate"); got != challenge { t.Errorf("WWW-Authenticate = %q, want %q", got, challenge) } }) } } // TestMetricsRecordsServedRequest checks that the metrics middleware // records a request it served, so /metrics reports it. It is the only test // in this package that sets up the metrics middleware, which registers with // the process-wide Prometheus registry and can do so only once. func TestMetricsRecordsServedRequest(t *testing.T) { t.Parallel() // The line /metrics shows once one GET /test has been served. const want = `http_request_duration_seconds_count{` + `code="200",handler="/test",method="GET",service=""} 1` mw := &Middleware{log: slog.Default(), config: &config.Config{}} handler := mw.Metrics()(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })) handler.ServeHTTP(httptest.NewRecorder(), httptest.NewRequestWithContext( t.Context(), http.MethodGet, "/test", nil)) rec := httptest.NewRecorder() promhttp.Handler().ServeHTTP(rec, httptest.NewRequestWithContext( t.Context(), http.MethodGet, "/metrics", nil)) if !strings.Contains(rec.Body.String(), want) { t.Errorf("/metrics does not report the GET /test served; "+ "want the line %q in:\n%s", want, rec.Body.String()) } } // TestLoggingLeavesOutSubmittedSigningKey checks that a login, a POST / // whose form carries the signing key, leaves no trace of the key in the // request's log line. func TestLoggingLeavesOutSubmittedSigningKey(t *testing.T) { t.Parallel() const signingKey = "test-signing-key-0123456789abcdef" var buf bytes.Buffer mw := newTestMiddleware(t, &buf) // The handler reads the key from the form, as the login handler does. handler := mw.Logging()(http.HandlerFunc( func(w http.ResponseWriter, r *http.Request) { if got := r.FormValue("key"); got != signingKey { t.Errorf("key in form = %q, want %q", got, signingKey) } w.WriteHeader(http.StatusSeeOther) })) form := url.Values{"key": {signingKey}} req := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") handler.ServeHTTP(httptest.NewRecorder(), req) if !strings.Contains(buf.String(), `"method":"POST"`) { t.Fatalf("no log line for the request; got %q", buf.String()) } if strings.Contains(buf.String(), signingKey) { t.Errorf("log output contains the signing key; got %q", buf.String()) } } func TestSecurityHeaders(t *testing.T) { t.Parallel() // Create middleware instance cfg := &config.Config{} mw := &Middleware{ log: slog.Default(), config: cfg, } // Create a test handler testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) // Wrap with security headers middleware handler := mw.SecurityHeaders()(testHandler) // Make a test request req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/test", nil) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) // Check security headers tests := []struct { header string want string }{ {"X-Content-Type-Options", "nosniff"}, {"X-Frame-Options", "DENY"}, {"Referrer-Policy", "strict-origin-when-cross-origin"}, {"X-XSS-Protection", "0"}, } for _, tt := range tests { t.Run(tt.header, func(t *testing.T) { t.Parallel() got := rec.Header().Get(tt.header) if got != tt.want { t.Errorf("%s = %q, want %q", tt.header, got, tt.want) } }) } } func TestSecurityHeaders_PolicyHeaders(t *testing.T) { t.Parallel() cfg := &config.Config{} mw := &Middleware{ log: slog.Default(), config: cfg, } testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) handler := mw.SecurityHeaders()(testHandler) req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/test", nil) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) // The login and generator pages load their script and stylesheet from // /static, so the policy allows no inline script or style. csp := rec.Header().Get("Content-Security-Policy") if strings.Contains(csp, "unsafe-inline") { t.Errorf("Content-Security-Policy allows unsafe-inline: %q", csp) } tests := []struct { header string want string }{ {"Strict-Transport-Security", "max-age=31536000; includeSubDomains"}, { "Content-Security-Policy", "default-src 'self'; " + "script-src 'self'; " + "style-src 'self'; " + "object-src 'none'; " + "base-uri 'self'; " + "form-action 'self'; " + "frame-ancestors 'none'", }, { "Permissions-Policy", "accelerometer=(), autoplay=(), camera=(), " + "display-capture=(), geolocation=(), gyroscope=(), " + "magnetometer=(), microphone=(), payment=(), usb=()", }, } for _, tt := range tests { t.Run(tt.header, func(t *testing.T) { t.Parallel() got := rec.Header().Get(tt.header) if got != tt.want { t.Errorf("%s = %q, want %q", tt.header, got, tt.want) } }) } } func TestSecurityHeaders_PreservesExistingHeaders(t *testing.T) { t.Parallel() cfg := &config.Config{} mw := &Middleware{ log: slog.Default(), config: cfg, } // Handler that sets its own headers testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("X-Custom-Header", "custom-value") w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) }) handler := mw.SecurityHeaders()(testHandler) req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/test", nil) rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) // Security headers should be present if rec.Header().Get("X-Content-Type-Options") != "nosniff" { t.Error("X-Content-Type-Options not set") } // Custom headers should still be there if rec.Header().Get("X-Custom-Header") != "custom-value" { t.Error("Custom header was overwritten") } if rec.Header().Get("Content-Type") != "application/json" { t.Error("Content-Type was overwritten") } }