Files
dnswatcher/internal/server/routes_test.go
T
sneak 879d41eb2b
check / check (push) Successful in 1m31s
metrics: rate limit /metrics per client address before Basic Auth (closes #101)
/metrics is behind a password, and REPO_POLICIES.md requires rate
limiting on password logins. Each client address may now send it 30
requests a minute, counted by httprate before Basic Auth, so failed
logins use up the allowance and a request over it gets 429 without
the password being checked. The address is the one the existing
trusted-proxy logic in internal/middleware works out, with IPv6
addresses grouped by /64; an IPv4 address a proxy reports in
IPv6-mapped form counts as the plain IPv4 address. A Prometheus
server scraping every 15 seconds sends 4 requests a minute.

Model: opus-5-5
2026-10-01 20:07:57 +00:00

291 lines
7.1 KiB
Go

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,
)
}
}
}
}
// metricsRequest builds a GET for /metrics from remoteAddr that logs
// in with the given password.
func metricsRequest(
t *testing.T,
remoteAddr string,
password string,
) *http.Request {
t.Helper()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, "/metrics", nil,
)
req.RemoteAddr = remoteAddr
req.SetBasicAuth(metricsUsername, password)
return req
}
// TestMetricsRateLimitComesBeforeAuth checks that failed logins to
// /metrics count against the rate limit; that once an address is over
// it, even the right password gets 429, with the same body as a wrong
// one; and that another address still gets in.
func TestMetricsRateLimitComesBeforeAuth(t *testing.T) {
viper.Reset()
t.Setenv("DNSWATCHER_TARGETS", "example.com")
t.Setenv("DNSWATCHER_METRICS_USERNAME", metricsUsername)
t.Setenv("DNSWATCHER_METRICS_PASSWORD", metricsPassword)
const (
guesser = "198.51.100.1:4000"
other = "198.51.100.2:4000"
// Far more guesses than the rate limit allows.
maxGuesses = 1000
)
srv := routedServer(t)
var guess *httptest.ResponseRecorder
for range maxGuesses {
guess = serve(srv, metricsRequest(t, guesser, "wrong"))
if guess.Code != http.StatusUnauthorized {
break
}
}
if guess.Code != http.StatusTooManyRequests {
t.Fatalf("wrong password: status = %d, want 429", guess.Code)
}
right := serve(srv, metricsRequest(t, guesser, metricsPassword))
if right.Code != http.StatusTooManyRequests {
t.Errorf("right password: status = %d, want 429", right.Code)
}
if right.Body.String() != guess.Body.String() {
t.Errorf(
"429 body with right password = %q, with wrong one = %q",
right.Body.String(), guess.Body.String(),
)
}
rec := serve(srv, metricsRequest(t, other, metricsPassword))
if rec.Code != http.StatusOK {
t.Errorf("another address: status = %d, want 200", rec.Code)
}
}