package middleware_test import ( "bytes" "log/slog" "net/http" "net/http/httptest" "net/netip" "strings" "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) } } } // TestMaxBodyBytesRejectsOversizeOnNonReadingRoute confirms the // limit is enforced even for a handler that never reads the body // (for example the health check), via the Content-Length check. func TestMaxBodyBytesRejectsOversizeOnNonReadingRoute(t *testing.T) { t.Parallel() const limit = 16 called := false handler := (&middleware.Middleware{}).MaxBodyBytes(limit)( http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { called = true }), ) rec := httptest.NewRecorder() req := httptest.NewRequest( http.MethodPost, "/.well-known/healthcheck", strings.NewReader(strings.Repeat("x", limit+1)), ) handler.ServeHTTP(rec, req) if rec.Code != http.StatusRequestEntityTooLarge { t.Fatalf("status = %d, want %d", rec.Code, http.StatusRequestEntityTooLarge) } if called { t.Fatal("handler ran despite oversize body") } } func TestMaxBodyBytesAllowsWithinLimit(t *testing.T) { t.Parallel() const limit = 64 handler := (&middleware.Middleware{}).MaxBodyBytes(limit)( http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }), ) rec := httptest.NewRecorder() req := httptest.NewRequest( http.MethodPost, "/api/v1/reports", strings.NewReader(`{"clientId":"c1"}`), ) handler.ServeHTTP(rec, req) if rec.Code != http.StatusOK { t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK) } } func TestRecovererReturns500AndLogsThroughSlog(t *testing.T) { t.Parallel() var logbuf bytes.Buffer mw := middleware.NewWithLogger( slog.New(slog.NewJSONHandler(&logbuf, nil)), ) handler := mw.Recoverer()( http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { panic("boom") }), ) rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/", http.NoBody) handler.ServeHTTP(rec, req) if rec.Code != http.StatusInternalServerError { t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError) } out := logbuf.String() if !strings.Contains(out, "panic recovered") { t.Fatalf("panic was not logged through slog: %q", out) } if !strings.Contains(out, `"level":"ERROR"`) { t.Fatalf("panic log was not structured JSON at error level: %q", out) } }