diff --git a/pkg/middleware/middleware.go b/pkg/middleware/middleware.go index 9cb41c474ad..cac1e5a4972 100644 --- a/pkg/middleware/middleware.go +++ b/pkg/middleware/middleware.go @@ -128,6 +128,16 @@ func WithSecurityHeaders(hdlr http.Handler) http.HandlerFunc { w.Header().Set("X-DNS-Prefetch-Control", "off") // Less information leakage about what domains we link to w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin") + w.Header().Set("Cache-Control", "no-cache, no-store, must-revalidate") + w.Header().Set("Pragma", "no-cache") + hdlr.ServeHTTP(w, r) + } +} + +func WithStaticCacheHeaders(hdlr http.Handler) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "public, max-age=31536000, immutable") + w.Header().Del("Pragma") hdlr.ServeHTTP(w, r) } } diff --git a/pkg/middleware/middleware_test.go b/pkg/middleware/middleware_test.go new file mode 100644 index 00000000000..12c3e3bf5bf --- /dev/null +++ b/pkg/middleware/middleware_test.go @@ -0,0 +1,73 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" +) + +func TestWithSecurityHeaders(t *testing.T) { + inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + handler := WithSecurityHeaders(inner) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + expected := map[string]string{ + "X-Content-Type-Options": "nosniff", + "X-Frame-Options": "DENY", + "X-DNS-Prefetch-Control": "off", + "Referrer-Policy": "strict-origin-when-cross-origin", + "Cache-Control": "no-cache, no-store, must-revalidate", + "Pragma": "no-cache", + } + + for header, want := range expected { + got := rr.Header().Get(header) + if got != want { + t.Errorf("header %s = %q, want %q", header, got, want) + } + } +} + +func TestWithStaticCacheHeaders(t *testing.T) { + inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + handler := WithSecurityHeaders(WithStaticCacheHeaders(inner)) + + req := httptest.NewRequest(http.MethodGet, "/static/main-bundle-abc123.min.js", nil) + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + wantCache := "public, max-age=31536000, immutable" + if got := rr.Header().Get("Cache-Control"); got != wantCache { + t.Errorf("Cache-Control = %q, want %q", got, wantCache) + } + if got := rr.Header().Get("Pragma"); got != "" { + t.Errorf("Pragma = %q, want empty (deleted)", got) + } +} + +func TestNonStaticEndpointHasNoCacheHeaders(t *testing.T) { + inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + handler := WithSecurityHeaders(inner) + + req := httptest.NewRequest(http.MethodGet, "/api/console/version", nil) + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + wantCache := "no-cache, no-store, must-revalidate" + if got := rr.Header().Get("Cache-Control"); got != wantCache { + t.Errorf("Cache-Control = %q, want %q", got, wantCache) + } + wantPragma := "no-cache" + if got := rr.Header().Get("Pragma"); got != wantPragma { + t.Errorf("Pragma = %q, want %q", got, wantPragma) + } +} diff --git a/pkg/server/server.go b/pkg/server/server.go index 2cdfae2f8c5..1852332b8d6 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -328,7 +328,7 @@ func (s *Server) HTTPHandler() (http.Handler, error) { handleFunc("/api/", notFoundHandler) staticHandler := http.StripPrefix(proxy.SingleJoiningSlash(s.BaseURL.Path, "/static/"), disableDirectoryListing(http.FileServer(http.Dir(s.PublicDir)))) - handle("/static/", middleware.WithGZIPEncoding(middleware.WithSecurityHeaders(staticHandler))) + handle("/static/", middleware.WithGZIPEncoding(middleware.WithStaticCacheHeaders(staticHandler))) // Register robots.txt at the origin root so crawlers can find it at /robots.txt // regardless of s.BaseURL.Path (e.g., /console/).