package middleware_test import ( "bytes" "encoding/json" "errors" "log/slog" "net/http" "net/http/httptest" "net/netip" "strings" "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) } } } // 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.NewRequestWithContext(t.Context(), 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") } if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" { t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n") } got := rec.Header().Get("Content-Type") if got != "application/json; charset=utf-8" { t.Errorf("Content-Type = %q, want a JSON content type", got) } } 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.NewRequestWithContext(t.Context(), 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.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody) handler.ServeHTTP(rec, req) if rec.Code != http.StatusInternalServerError { t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError) } var record map[string]any err := json.Unmarshal(logbuf.Bytes(), &record) if err != nil { t.Fatalf("panic log is not one JSON record: %v (%q)", err, logbuf.String()) } if record["msg"] != "panic recovered" || record["level"] != "ERROR" { t.Errorf("log record = %v, want msg %q at level ERROR", record, "panic recovered") } if record["panic"] != "boom" { t.Errorf("panic field = %v, want %q", record["panic"], "boom") } stack, _ := record["stack"].(string) if !strings.HasPrefix(stack, "goroutine ") { t.Errorf("stack field = %q, want a stack trace", stack) } } // TestRecovererRepanicsOnAbortHandler checks that a handler aborting // with http.ErrAbortHandler is not treated as a crash: Recoverer // panics again so the server aborts the response, and logs nothing. func TestRecovererRepanicsOnAbortHandler(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(http.ErrAbortHandler) }), ) rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody) var recovered any func() { defer func() { recovered = recover() }() handler.ServeHTTP(rec, req) }() err, _ := recovered.(error) if !errors.Is(err, http.ErrAbortHandler) { t.Errorf("Recoverer panicked with %v, want http.ErrAbortHandler", recovered) } if logbuf.Len() != 0 { t.Errorf("abort was logged: %q", logbuf.String()) } }