package middleware_test import ( "io" "net/http" "net/http/httptest" "strings" "testing" "time" "sneak.berlin/go/simplexcalc/internal/config" "sneak.berlin/go/simplexcalc/internal/globals" "sneak.berlin/go/simplexcalc/internal/logger" "sneak.berlin/go/simplexcalc/internal/middleware" "sneak.berlin/go/simplexcalc/internal/telemetry" ) // newMiddleware builds the set against a given config, with logging // discarded and telemetry disabled. func newMiddleware(t *testing.T, cfg *config.Config) *middleware.Middleware { t.Helper() g := &globals.Globals{Appname: "simplexcalc", Version: "test"} log, err := logger.New(nil, logger.Params{Globals: g, Output: io.Discard}) if err != nil { t.Fatalf("building logger: %v", err) } sentry, err := telemetry.NewSentry(nil, telemetry.SentryParams{ Config: cfg, Globals: g, Logger: log, }) if err != nil { t.Fatalf("building sentry: %v", err) } metrics, err := telemetry.NewMetrics(telemetry.MetricsParams{Config: cfg}) if err != nil { t.Fatalf("building metrics: %v", err) } mw, err := middleware.New(nil, middleware.Params{ Config: cfg, Logger: log, Sentry: sentry, Metrics: metrics, }) if err != nil { t.Fatalf("building middleware: %v", err) } return mw } // getReq and postReq build requests carrying the test's context, so a // handler that respects cancellation is exercised the way the server // exercises it. func getReq(t *testing.T) *http.Request { t.Helper() return httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil) } func postReq(t *testing.T, body string) *http.Request { t.Helper() return httptest.NewRequestWithContext( t.Context(), http.MethodPost, "/", strings.NewReader(body), ) } func testConfig() *config.Config { return &config.Config{ Port: 8080, HSTS: true, MaxRequestBody: 1024, RequestTimeout: time.Second, ShutdownGrace: time.Second, CSRFKeyEphemeral: true, } } // TestRecovererAnswers500 is the point of the panic middleware: net/http // on its own drops the connection, which tells the client nothing about // whose fault it was. func TestRecovererAnswers500(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) h := mw.Recoverer()(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { panic("boom") })) w := httptest.NewRecorder() h.ServeHTTP(w, getReq(t)) if w.Code != http.StatusInternalServerError { t.Fatalf("status = %d, want 500", w.Code) } if w.Body.Len() == 0 { t.Error("a 500 with no body tells the client nothing") } // The panic value must not reach the client. if strings.Contains(w.Body.String(), "boom") { t.Error("the panic value was leaked in the response body") } } // TestRecovererPassesThroughSuccess: the recovery wrapper must be // invisible when nothing goes wrong, including for the response body. func TestRecovererPassesThroughSuccess(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) h := mw.Recoverer()(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusTeapot) _, _ = w.Write([]byte("fine")) })) w := httptest.NewRecorder() h.ServeHTTP(w, getReq(t)) if w.Code != http.StatusTeapot || w.Body.String() != "fine" { t.Errorf("status = %d body = %q", w.Code, w.Body.String()) } } // TestSecurityHeadersOnEveryResponse, including responses the handler // never got to write. func TestSecurityHeadersOnEveryResponse(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) notFound := func(w http.ResponseWriter, _ *http.Request) { http.Error(w, "not found", http.StatusNotFound) } h := mw.SecurityHeaders()(http.HandlerFunc(notFound)) w := httptest.NewRecorder() h.ServeHTTP(w, getReq(t)) want := map[string]string{ "X-Frame-Options": "DENY", "X-Content-Type-Options": "nosniff", "Referrer-Policy": "strict-origin-when-cross-origin", "Strict-Transport-Security": "max-age=31536000; includeSubDomains", } for header, value := range want { if got := w.Header().Get(header); got != value { t.Errorf("%s = %q, want %q", header, got, value) } } csp := w.Header().Get("Content-Security-Policy") if !strings.Contains(csp, "default-src 'self'") { t.Errorf("CSP = %q", csp) } if strings.Contains(csp, "unsafe-inline") { t.Error("the CSP permits inline script or style") } } // TestHSTSOffWhenDisabled: the header must be absent, not empty, so a // developer's browser is never pinned to HTTPS on localhost. func TestHSTSOffWhenDisabled(t *testing.T) { t.Parallel() cfg := testConfig() cfg.HSTS = false mw := newMiddleware(t, cfg) noop := func(_ http.ResponseWriter, _ *http.Request) {} h := mw.SecurityHeaders()(http.HandlerFunc(noop)) w := httptest.NewRecorder() h.ServeHTTP(w, getReq(t)) if _, ok := w.Header()["Strict-Transport-Security"]; ok { t.Error("HSTS was sent with HSTS disabled") } } // TestBodyLimitRefusesDeclaredOversize: a Content-Length over the cap is // refused before the body transfers. func TestBodyLimitRefusesDeclaredOversize(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) reached := false h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { reached = true })) req := postReq(t, strings.Repeat("x", 2048)) w := httptest.NewRecorder() h.ServeHTTP(w, req) if w.Code != http.StatusRequestEntityTooLarge { t.Errorf("status = %d, want 413", w.Code) } if reached { t.Error("the handler ran for an oversized request") } } // TestBodyLimitCapsUndeclaredBody is the case Content-Length cannot // catch: a body that arrives without one, or with a lying one, must // still fail at the cap rather than being read in full. func TestBodyLimitCapsUndeclaredBody(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) var readErr error h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { _, readErr = io.ReadAll(r.Body) })) req := postReq(t, strings.Repeat("x", 4096)) // Undeclared length: what a chunked upload looks like here. req.ContentLength = -1 h.ServeHTTP(httptest.NewRecorder(), req) if readErr == nil { t.Error("reading past the cap succeeded; the limit is not enforced on the read") } } // TestBodyLimitAllowsNormalRequests, so the cap is not just "refuse // everything". func TestBodyLimitAllowsNormalRequests(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) got := "" h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { b, _ := io.ReadAll(r.Body) got = string(b) })) req := postReq(t, "small") h.ServeHTTP(httptest.NewRecorder(), req) if got != "small" { t.Errorf("body = %q, want %q", got, "small") } } // TestRequestIDIsAssignedAndNotBorrowed: the id must be this process's, // so a client cannot collide two unrelated requests in the log. func TestRequestIDIsAssignedAndNotBorrowed(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) var inHandler string h := mw.RequestID()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { inHandler = middleware.RequestIDFrom(r.Context()) })) req := getReq(t) req.Header.Set(middleware.RequestIDHeader, "client-supplied") w := httptest.NewRecorder() h.ServeHTTP(w, req) if inHandler == "" { t.Fatal("no request id reached the handler") } if inHandler == "client-supplied" { t.Error("the client's request id was trusted") } if w.Header().Get(middleware.RequestIDHeader) != inHandler { t.Error("the response header does not carry the id the handler saw") } } // TestTimeoutGivesHandlerADeadline. The handler is what has to respect // it, so what is asserted here is that the deadline is there at all. func TestTimeoutGivesHandlerADeadline(t *testing.T) { t.Parallel() mw := newMiddleware(t, testConfig()) var hasDeadline bool h := mw.Timeout()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { _, hasDeadline = r.Context().Deadline() })) h.ServeHTTP(httptest.NewRecorder(), getReq(t)) if !hasDeadline { t.Error("the handler's context carries no deadline") } }