From d9ac325ea0fb8ba01c5d59e4203bc422a6f96c0c Mon Sep 17 00:00:00 2001 From: clawbot <35+clawbot@noreply.example.org> Date: Thu, 1 Oct 2026 18:02:35 +0000 Subject: [PATCH] server: limit wildcard CORS to the public routes (closes #100) The CORS wildcard was global, so it also covered the Basic-Auth protected /metrics, which REPO_POLICIES.md forbids, and it allowed POST, PUT and DELETE, which no route serves, plus the Authorization and X-CSRF-Token headers. CORS now sits on a router holding only the public routes and allows GET and OPTIONS with the Accept and Content-Type headers. /metrics gets no CORS at all. Both are mounted routers rather than a Group: chi answers OPTIONS on a Group's route with 405 before its middleware runs, and any method /metrics does not register would otherwise fall through to the public router. So every method on /metrics now meets Basic Auth first, and /metrics/ is served like /metrics. Model: opus-5-5 --- TODO.md | 2 + internal/middleware/middleware.go | 15 +- internal/server/routes.go | 37 +++-- internal/server/routes_test.go | 221 ++++++++++++++++++++++++++++++ 4 files changed, 252 insertions(+), 23 deletions(-) create mode 100644 internal/server/routes_test.go diff --git a/TODO.md b/TODO.md index a7c1ce8..23a1f8f 100644 --- a/TODO.md +++ b/TODO.md @@ -19,6 +19,8 @@ Rationale, Design, TODO, License, Author) if any are still missing. # Completed Steps +- 2026-10-01: wildcard CORS now applies only to the public routes, not to + `/metrics`, and allows only the methods they serve (closes #100). - 2026-10-01: `internal/state` and `internal/watcher` no longer export test-only constructors: two moved to `export_test.go`, one is deleted (closes #111). - 2026-10-01: notify shutdown tests use one timing constant per meaning, name diff --git a/internal/middleware/middleware.go b/internal/middleware/middleware.go index 03f435e..4ba05a8 100644 --- a/internal/middleware/middleware.go +++ b/internal/middleware/middleware.go @@ -223,17 +223,14 @@ func realIP(r *http.Request) string { return addr } -// CORS returns CORS middleware. +// CORS returns middleware that lets any origin read a response. It is +// for the public, read-only routes only, so it allows only the +// methods those routes serve and no Authorization header. func (m *Middleware) CORS() func(http.Handler) http.Handler { return cors.Handler(cors.Options{ - AllowedOrigins: []string{"*"}, - AllowedMethods: []string{ - "GET", "POST", "PUT", "DELETE", "OPTIONS", - }, - AllowedHeaders: []string{ - "Accept", "Authorization", - "Content-Type", "X-CSRF-Token", - }, + AllowedOrigins: []string{"*"}, + AllowedMethods: []string{"GET", "OPTIONS"}, + AllowedHeaders: []string{"Accept", "Content-Type"}, ExposedHeaders: []string{"Link"}, AllowCredentials: false, MaxAge: corsMaxAge, diff --git a/internal/server/routes.go b/internal/server/routes.go index 5c71d84..6305e1d 100644 --- a/internal/server/routes.go +++ b/internal/server/routes.go @@ -23,14 +23,20 @@ func (s *Server) SetupRoutes() { s.router.Use(chimw.RequestID) s.router.Use(s.mw.SecurityHeaders()) s.router.Use(s.mw.Logging()) - s.router.Use(s.mw.CORS()) s.router.Use(chimw.Timeout(requestTimeout)) + // Public, unauthenticated, read-only routes, the only ones + // REPO_POLICIES.md allows wildcard CORS on. CORS is middleware of + // this whole router, not of a Group, so that it also answers + // OPTIONS preflight requests, which no route here registers. + public := chi.NewRouter() + public.Use(s.mw.CORS()) + // Dashboard (read-only web UI) - s.router.Get("/", s.handlers.HandleDashboard()) + public.Get("/", s.handlers.HandleDashboard()) // Static assets (embedded CSS/JS) - s.router.Mount( + public.Mount( "/s", http.StripPrefix( "/s", @@ -39,27 +45,30 @@ func (s *Server) SetupRoutes() { ) // Health check (standard well-known path) - s.router.Get( + public.Get( "/.well-known/healthcheck", s.handlers.HandleHealthCheck(), ) // Legacy health check (keep for backward compatibility) - s.router.Get("/health", s.handlers.HandleHealthCheck()) + public.Get("/health", s.handlers.HandleHealthCheck()) // API v1 routes - s.router.Route("/api/v1", func(r chi.Router) { + public.Route("/api/v1", func(r chi.Router) { r.Get("/status", s.handlers.HandleStatus()) }) - // Metrics endpoint (optional, with basic auth) + s.router.Mount("/", public) + + // Metrics endpoint (optional, with basic auth) and no CORS: a + // Prometheus scraper is not a browser. It is mounted rather than + // added with Get so that every method on /metrics, OPTIONS + // included, ends here instead of falling through to the public + // router and its CORS. if s.params.Config.MetricsUsername != "" { - s.router.Group(func(r chi.Router) { - r.Use(s.mw.MetricsAuth()) - r.Get( - "/metrics", - promhttp.Handler().ServeHTTP, - ) - }) + metrics := chi.NewRouter() + metrics.Use(s.mw.MetricsAuth()) + metrics.Get("/", promhttp.Handler().ServeHTTP) + s.router.Mount("/metrics", metrics) } } diff --git a/internal/server/routes_test.go b/internal/server/routes_test.go new file mode 100644 index 0000000..a16c762 --- /dev/null +++ b/internal/server/routes_test.go @@ -0,0 +1,221 @@ +package server_test + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/spf13/viper" + + "sneak.berlin/go/dnswatcher/internal/server" +) + +// Credentials for /metrics, which is only routed when a username is set. +const ( + metricsUsername = "scraper" + metricsPassword = "scrape-secret" +) + +// The tests below set env vars and touch viper global state, so like +// the config tests they cannot use t.Parallel. + +// routedServer builds the server with its routes set up, ready to serve +// test requests. The caller must first configure viper. +func routedServer(t *testing.T) *server.Server { + t.Helper() + + srv := buildServer(t) + srv.SetupRoutes() + + return srv +} + +// crossOriginRequest builds a request as a browser sends it from a page +// on another site. +func crossOriginRequest( + t *testing.T, + method string, + target string, +) *http.Request { + t.Helper() + + req := httptest.NewRequestWithContext(t.Context(), method, target, nil) + req.Header.Set("Origin", "https://example.net") + + return req +} + +// preflightRequest builds the OPTIONS request a browser sends before a +// cross-origin request with the given method and request headers. +func preflightRequest( + t *testing.T, + target string, + method string, + headers string, +) *http.Request { + t.Helper() + + req := crossOriginRequest(t, http.MethodOptions, target) + req.Header.Set("Access-Control-Request-Method", method) + + if headers != "" { + req.Header.Set("Access-Control-Request-Headers", headers) + } + + return req +} + +func serve( + srv *server.Server, + req *http.Request, +) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + + return rec +} + +// publicPaths returns one path on each public route. +func publicPaths() []string { + return []string{ + "/", + "/s/css/tailwind.min.css", + "/api/v1/status", + "/health", + "/.well-known/healthcheck", + } +} + +// TestPublicRoutesAllowAnyOrigin checks that every public route answers +// a cross-origin GET with the CORS wildcard. +func TestPublicRoutesAllowAnyOrigin(t *testing.T) { + viper.Reset() + t.Setenv("DNSWATCHER_TARGETS", "example.com") + t.Setenv("DNSWATCHER_METRICS_USERNAME", metricsUsername) + t.Setenv("DNSWATCHER_METRICS_PASSWORD", metricsPassword) + + srv := routedServer(t) + + for _, path := range publicPaths() { + rec := serve(srv, crossOriginRequest(t, http.MethodGet, path)) + + if rec.Code != http.StatusOK { + t.Errorf("GET %s: status = %d, want 200", path, rec.Code) + } + + got := rec.Header().Get("Access-Control-Allow-Origin") + if got != "*" { + t.Errorf( + "GET %s: Access-Control-Allow-Origin = %q, want %q", + path, got, "*", + ) + } + } +} + +// TestMetricsHasNoCORS checks that no request to the Basic-Auth +// protected /metrics, preflight included, gets a CORS header. +func TestMetricsHasNoCORS(t *testing.T) { + viper.Reset() + t.Setenv("DNSWATCHER_TARGETS", "example.com") + t.Setenv("DNSWATCHER_METRICS_USERNAME", metricsUsername) + t.Setenv("DNSWATCHER_METRICS_PASSWORD", metricsPassword) + + srv := routedServer(t) + + authenticated := crossOriginRequest(t, http.MethodGet, "/metrics") + authenticated.SetBasicAuth(metricsUsername, metricsPassword) + + tests := []struct { + name string + req *http.Request + wantStatus int + }{ + { + "authenticated GET", + authenticated, + http.StatusOK, + }, + { + "unauthenticated GET", + crossOriginRequest(t, http.MethodGet, "/metrics"), + http.StatusUnauthorized, + }, + { + "preflight", + preflightRequest(t, "/metrics", http.MethodGet, ""), + http.StatusUnauthorized, + }, + } + + for _, tt := range tests { + rec := serve(srv, tt.req) + + if rec.Code != tt.wantStatus { + t.Errorf( + "%s: status = %d, want %d", + tt.name, rec.Code, tt.wantStatus, + ) + } + + got := rec.Header().Get("Access-Control-Allow-Origin") + if got != "" { + t.Errorf( + "%s: Access-Control-Allow-Origin = %q, want none", + tt.name, got, + ) + } + } +} + +// TestPreflightAllowsOnlyWhatPublicRoutesServe checks what each public +// route agrees to in a CORS preflight: GET, but not POST, PUT or +// DELETE, which no route serves, and not the Authorization or +// X-CSRF-Token headers, which no public route reads. It checks every +// public route because one added with Get, such as /health, answers a +// preflight only while CORS is middleware of a whole router; in a +// Group, chi would answer it with 405 and no CORS headers. +func TestPreflightAllowsOnlyWhatPublicRoutesServe(t *testing.T) { + viper.Reset() + t.Setenv("DNSWATCHER_TARGETS", "example.com") + t.Setenv("DNSWATCHER_METRICS_USERNAME", metricsUsername) + t.Setenv("DNSWATCHER_METRICS_PASSWORD", metricsPassword) + + srv := routedServer(t) + + tests := []struct { + method string + headers string + allowed bool + }{ + {http.MethodGet, "", true}, + {http.MethodGet, "Content-Type", true}, + {http.MethodPost, "", false}, + {http.MethodPut, "", false}, + {http.MethodDelete, "", false}, + {http.MethodGet, "Authorization", false}, + {http.MethodGet, "X-CSRF-Token", false}, + } + + for _, path := range publicPaths() { + for _, tt := range tests { + rec := serve(srv, preflightRequest( + t, path, tt.method, tt.headers, + )) + + want := "" + if tt.allowed { + want = tt.method + } + + got := rec.Header().Get("Access-Control-Allow-Methods") + if got != want { + t.Errorf( + "preflight to %s for %s with headers %q: "+ + "Access-Control-Allow-Methods = %q, want %q", + path, tt.method, tt.headers, got, want, + ) + } + } + } +} -- 2.54.0