check / check (push) Successful in 50s
The request log wrote the URL, User-Agent, Referer and other request-supplied strings with no length limit, and the server accepts headers up to 1 MiB, so one request could put about 1 MiB per field into a log line. Every string the request log takes from the request, including the request ID chi copies from X-Request-Id, is now cut to the 128-byte bound the report handler already used. That bound and its helper moved from the handlers package to the logger package so both use the one copy. Model: opus-5-5
555 lines
15 KiB
Go
555 lines
15 KiB
Go
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/logger"
|
|
"sneak.berlin/go/netwatch/internal/middleware"
|
|
|
|
chimiddleware "github.com/go-chi/chi/v5/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 || !strings.Contains(err.Error(), "TRUSTED_PROXIES") {
|
|
t.Fatalf("error = %v, want one naming TRUSTED_PROXIES", err)
|
|
}
|
|
}
|
|
|
|
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())
|
|
}
|
|
}
|
|
|
|
// TestLoggingCutsRequestStringsToBound sends an over-long URL and
|
|
// over-long header values, and checks the request log writes each
|
|
// one cut to logger.MaxLoggedFieldBytes.
|
|
func TestLoggingCutsRequestStringsToBound(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
long := strings.Repeat("a", 2*logger.MaxLoggedFieldBytes)
|
|
|
|
var logbuf bytes.Buffer
|
|
|
|
mw := middleware.NewWithLogger(
|
|
slog.New(slog.NewJSONHandler(&logbuf, nil)),
|
|
)
|
|
|
|
handler := chimiddleware.RequestID(mw.Logging()(okHandler()))
|
|
|
|
req := httptest.NewRequestWithContext(t.Context(),
|
|
http.MethodGet, "/"+long, http.NoBody)
|
|
req.Header.Set("User-Agent", long)
|
|
req.Header.Set("Referer", long)
|
|
req.Header.Set("X-Request-Id", long)
|
|
|
|
handler.ServeHTTP(httptest.NewRecorder(), req)
|
|
|
|
var logged map[string]any
|
|
|
|
err := json.Unmarshal(logbuf.Bytes(), &logged)
|
|
if err != nil {
|
|
t.Fatalf("log line not JSON: %v (%q)", err, logbuf.String())
|
|
}
|
|
|
|
want := map[string]string{
|
|
"url": ("/" + long)[:logger.MaxLoggedFieldBytes],
|
|
"useragent": long[:logger.MaxLoggedFieldBytes],
|
|
"referer": long[:logger.MaxLoggedFieldBytes],
|
|
"request_id": long[:logger.MaxLoggedFieldBytes],
|
|
}
|
|
|
|
for field, value := range want {
|
|
if logged[field] != value {
|
|
t.Errorf("logged %s = %q, want it cut to %d bytes",
|
|
field, logged[field], logger.MaxLoggedFieldBytes)
|
|
}
|
|
}
|
|
}
|
|
|
|
// 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)
|
|
}
|
|
})
|
|
}
|
|
}
|