package middleware_test import ( "net/http" "net/http/httptest" "net/netip" "testing" "sneak.berlin/go/netwatch/internal/middleware" ) const ( // loopbackPeer is a remote address inside the trusted-proxy allowlist. loopbackPeer = "127.0.0.1:5000" // forwardedIP is the client address presented via X-Forwarded-For. forwardedIP = "203.0.113.7" // realIP is the client address presented via X-Real-IP. realIP = "203.0.113.9" ) func mustPrefixes(t *testing.T, cidrs ...string) []netip.Prefix { t.Helper() prefixes, err := middleware.ParseTrustedProxies(cidrs) if err != nil { t.Fatalf("ParseTrustedProxies(%v): %v", cidrs, err) } return prefixes } func TestParseTrustedProxiesRejectsMalformed(t *testing.T) { t.Parallel() _, err := middleware.ParseTrustedProxies([]string{"not-a-cidr"}) if err == nil { t.Fatal("expected error for malformed CIDR, got nil") } } type clientIPCase struct { name string remoteAddr string xff string xRealIP string want string } func clientIPCases() []clientIPCase { return []clientIPCase{ { name: "trusted proxy uses forwarded-for", remoteAddr: loopbackPeer, xff: forwardedIP, want: forwardedIP, }, { name: "trusted proxy uses left-most of chain", remoteAddr: "10.1.2.3:5000", xff: forwardedIP + ", 10.1.2.3", want: forwardedIP, }, { name: "trusted proxy falls back to x-real-ip", remoteAddr: loopbackPeer, xRealIP: realIP, want: realIP, }, { name: "untrusted peer ignores forwarded-for", remoteAddr: "198.51.100.4:5000", xff: forwardedIP, want: "198.51.100.4", }, { name: "untrusted peer ignores x-real-ip", remoteAddr: "198.51.100.4:5000", xRealIP: realIP, want: "198.51.100.4", }, { name: "trusted proxy with no headers uses peer", remoteAddr: "10.1.2.3:5000", want: "10.1.2.3", }, { name: "trusted proxy with garbage header uses peer", remoteAddr: loopbackPeer, xff: "not-an-ip", want: "127.0.0.1", }, } } func TestClientIP(t *testing.T) { t.Parallel() trusted := mustPrefixes(t, "127.0.0.1/32", "::1/128", "10.0.0.0/8") for _, tc := range clientIPCases() { t.Run(tc.name, func(t *testing.T) { t.Parallel() header := http.Header{} if tc.xff != "" { header.Set("X-Forwarded-For", tc.xff) } if tc.xRealIP != "" { header.Set("X-Real-IP", tc.xRealIP) } got := middleware.ClientIP(tc.remoteAddr, header, trusted) if got != tc.want { t.Errorf("ClientIP() = %q, want %q", got, tc.want) } }) } } func TestSecurityHeaders(t *testing.T) { t.Parallel() handler := (&middleware.Middleware{}).SecurityHeaders()( http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }), ) rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody) handler.ServeHTTP(rec, req) want := map[string]string{ "Strict-Transport-Security": "max-age=31536000; includeSubDomains", "Content-Security-Policy": "default-src 'none'; frame-ancestors 'none'", "X-Frame-Options": "DENY", "X-Content-Type-Options": "nosniff", "Referrer-Policy": "no-referrer", "Permissions-Policy": "camera=(), microphone=(), geolocation=()", } for name, value := range want { if got := rec.Header().Get(name); got != value { t.Errorf("header %s = %q, want %q", name, got, value) } } }