Files
netwatch/backend/internal/middleware/middleware_test.go
T
clawbot 503399e020
check / check (push) Successful in 10s
fix(backend): report ingest correctness: 500 on a refused report, 413 on oversize, global body cap (closes #23)
A report the buffer refuses now returns 500 instead of a false `ok`.
Reports reach disk later, so a failed disk write is still answered 200
and shows in the log, and at shutdown as a failed stop with a non-zero
exit. An over-limit body returns 413; malformed JSON stays 400. A
MaxBodyBytes middleware caps every route at 1 MiB; a route group can
only lower that limit. The raw geo blob is no longer logged; client_id,
timestamp and decode error text are cut to 128 bytes before logging.
Panic recovery logs the panic value and stack through slog.

Model: opus-5-5
2026-09-28 20:39:39 +02:00

303 lines
7.3 KiB
Go

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())
}
}