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