package middleware_test import ( "bytes" "encoding/json" "errors" "log/slog" "net/http" "net/http/httptest" "net/netip" "strings" "testing" "testing/synctest" "time" "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 } // TestParseTrustedProxiesRejectsMalformed includes entries nginx would // read as another address or look up as a hostname, in the CIDR form // bin/entrypoint.sh gives "netwatch-server check-cidr". func TestParseTrustedProxiesRejectsMalformed(t *testing.T) { t.Parallel() for _, cidr := range []string{ "not-a-cidr", "10.0.0.1", "1.2.3/32", "172.30/32", "10/32", "cafe/32", "999.1.1.1/32", "10.0.0.0/33", "::1/129", "fe80::1%eth0/128", } { _, err := middleware.ParseTrustedProxies([]string{cidr}) if err == nil || !strings.Contains(err.Error(), "TRUSTED_PROXIES") { t.Errorf("%q: error = %v, want one naming TRUSTED_PROXIES", cidr, err) } } } func TestParseTrustedProxiesAcceptsCIDRs(t *testing.T) { t.Parallel() mustPrefixes(t, "172.17.0.1/32", "10.0.0.0/8", "2001:db8::1/128", "2001:db8::/32", "::ffff:192.0.2.1/128") } 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()) } } // okHandler stands in for the route a middleware guards. func okHandler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) } // TestRateLimitRefusesPastAllowanceThenResets checks one client // address: it may use its whole allowance at once, the next request // is refused with 429, and later it may send again. func TestRateLimitRefusesPastAllowanceThenResets(t *testing.T) { t.Parallel() // synctest runs this on a fake clock: time.Sleep returns at once, // with the clock moved on. synctest.Test(t, func(t *testing.T) { const perMinute = 2 handler := (&middleware.Middleware{}).RateLimit(perMinute)(okHandler()) post := func() *httptest.ResponseRecorder { rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/api/v1/reports", http.NoBody) handler.ServeHTTP(rec, req) return rec } for i := range perMinute { if code := post().Code; code != http.StatusOK { t.Fatalf("request %d: status = %d, want %d", i+1, code, http.StatusOK) } } rec := post() if rec.Code != http.StatusTooManyRequests { t.Fatalf("request past the allowance: status = %d, want %d", rec.Code, http.StatusTooManyRequests) } if got := rec.Body.String(); got != "{\"status\":\"error\"}\n" { t.Errorf("body = %q, want %q", got, "{\"status\":\"error\"}\n") } if got := rec.Header().Get("Retry-After"); got != "60" { t.Fatalf("Retry-After = %q, want %q", got, "60") } // httprate also counts the previous minute's requests, fading // them out over the current one, so two minutes on the whole // allowance is back. time.Sleep(2 * time.Minute) for i := range perMinute { if code := post().Code; code != http.StatusOK { t.Fatalf("two minutes later, request %d: status = %d, want %d", i+1, code, http.StatusOK) } } }) } // postForwarded sends handler a report from peer that names client in // X-Forwarded-For, and returns the status. func postForwarded( t *testing.T, handler http.Handler, peer, client string, ) int { t.Helper() rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/api/v1/reports", http.NoBody) req.RemoteAddr = peer req.Header.Set("X-Forwarded-For", client) handler.ServeHTTP(rec, req) return rec.Code } // TestRateLimitIsPerForwardedClient checks that clients behind a // trusted proxy each get their own allowance: the limit is keyed on // the client address clientIP resolves, not on the proxy's. func TestRateLimitIsPerForwardedClient(t *testing.T) { t.Parallel() const otherClient = "203.0.113.8" mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32")) handler := mw.RateLimit(1)(okHandler()) code := postForwarded(t, handler, loopbackPeer, forwardedIP) if code != http.StatusOK { t.Fatalf("first request: status = %d, want %d", code, http.StatusOK) } code = postForwarded(t, handler, loopbackPeer, forwardedIP) if code != http.StatusTooManyRequests { t.Fatalf("same client again: status = %d, want %d", code, http.StatusTooManyRequests) } code = postForwarded(t, handler, loopbackPeer, otherClient) if code != http.StatusOK { t.Fatalf("other client behind the same proxy: status = %d, want %d", code, http.StatusOK) } } // TestRateLimitIgnoresForwardedForFromUntrustedPeer checks that a // peer that is not a trusted proxy cannot get a fresh allowance by // naming a different client in X-Forwarded-For on each request. func TestRateLimitIgnoresForwardedForFromUntrustedPeer(t *testing.T) { t.Parallel() const untrustedPeer = "198.51.100.4:5000" mw := middleware.NewWithTrustedProxies(mustPrefixes(t, "127.0.0.1/32")) handler := mw.RateLimit(1)(okHandler()) code := postForwarded(t, handler, untrustedPeer, "203.0.113.8") if code != http.StatusOK { t.Fatalf("first request: status = %d, want %d", code, http.StatusOK) } code = postForwarded(t, handler, untrustedPeer, "203.0.113.9") if code != http.StatusTooManyRequests { t.Fatalf("same peer naming another client: status = %d, want %d", code, http.StatusTooManyRequests) } } // preflight sends cors the preflight request a browser makes before // it POSTs JSON from origin. func preflight( t *testing.T, cors func(http.Handler) http.Handler, origin string, ) *httptest.ResponseRecorder { t.Helper() rec := httptest.NewRecorder() req := httptest.NewRequestWithContext(t.Context(), http.MethodOptions, "/api/v1/reports", http.NoBody) req.Header.Set("Origin", origin) req.Header.Set("Access-Control-Request-Method", http.MethodPost) req.Header.Set("Access-Control-Request-Headers", "content-type") cors(okHandler()).ServeHTTP(rec, req) return rec } // TestCORSWithoutOriginsAddsNoHeaders checks the default: with no // origins configured, no origin is given any CORS header. func TestCORSWithoutOriginsAddsNoHeaders(t *testing.T) { t.Parallel() rec := preflight(t, (&middleware.Middleware{}).CORS(nil), "https://elsewhere.example") for name := range rec.Header() { if strings.HasPrefix(name, "Access-Control-") { t.Errorf("CORS header %s set with no origins configured", name) } } } func TestCORSAllowsOnlyListedOrigins(t *testing.T) { t.Parallel() const listed = "https://netwatch.example" cors := (&middleware.Middleware{}).CORS([]string{listed}) cases := []struct { name string origin string want string }{ {name: "listed origin allowed", origin: listed, want: listed}, {name: "other origin refused", origin: "https://elsewhere.example"}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { t.Parallel() rec := preflight(t, cors, tc.origin) got := rec.Header().Get("Access-Control-Allow-Origin") if got != tc.want { t.Errorf("Access-Control-Allow-Origin = %q, want %q", got, tc.want) } }) } }