package middleware_test import ( "net/http" "net/http/httptest" "net/netip" "testing" "sneak.berlin/go/netwatch/internal/middleware" ) 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: "127.0.0.1:5000", xff: "203.0.113.7", want: "203.0.113.7", }, { name: "trusted proxy uses left-most of chain", remoteAddr: "10.1.2.3:5000", xff: "203.0.113.7, 10.1.2.3", want: "203.0.113.7", }, { name: "trusted proxy falls back to x-real-ip", remoteAddr: "127.0.0.1:5000", xRealIP: "203.0.113.9", want: "203.0.113.9", }, { name: "untrusted peer ignores forwarded-for", remoteAddr: "198.51.100.4:5000", xff: "203.0.113.7", want: "198.51.100.4", }, { name: "untrusted peer ignores x-real-ip", remoteAddr: "198.51.100.4:5000", xRealIP: "203.0.113.9", 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: "127.0.0.1:5000", 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.NewRequest(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) } } }