metrics: rate limit /metrics per client address before Basic Auth (closes #101)
check / check (push) Successful in 1m7s
check / check (push) Successful in 1m7s
/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
This commit is contained in:
@@ -0,0 +1,10 @@
|
||||
package middleware
|
||||
|
||||
import "time"
|
||||
|
||||
// The /metrics rate limit, exported so the tests can count requests
|
||||
// against it.
|
||||
const (
|
||||
MetricsRequestLimit = metricsRequestLimit
|
||||
MetricsRequestWindow time.Duration = metricsRequestWindow
|
||||
)
|
||||
@@ -5,12 +5,14 @@ import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/99designs/basicauth-go"
|
||||
"github.com/go-chi/chi/v5/middleware"
|
||||
"github.com/go-chi/cors"
|
||||
"github.com/go-chi/httprate"
|
||||
"go.uber.org/fx"
|
||||
|
||||
"sneak.berlin/go/dnswatcher/internal/config"
|
||||
@@ -21,6 +23,17 @@ import (
|
||||
// corsMaxAge is the maximum age for CORS preflight responses.
|
||||
const corsMaxAge = 300
|
||||
|
||||
// Rate limit for /metrics: each client address may send
|
||||
// metricsRequestLimit requests per metricsRequestWindow. Every request
|
||||
// counts, so password guessing gets at most 30 tries a minute per
|
||||
// address. One Prometheus server scraping every 15 seconds sends 4
|
||||
// requests a minute, and two scraping every 5 seconds from one address
|
||||
// send 24, so normal scraping stays under the limit.
|
||||
const (
|
||||
metricsRequestLimit = 30
|
||||
metricsRequestWindow = time.Minute
|
||||
)
|
||||
|
||||
// Security response header values applied to every response.
|
||||
//
|
||||
// The CSP is as strict as the dashboard allows: the template ships no
|
||||
@@ -268,6 +281,32 @@ func (m *Middleware) SecurityHeaders() func(http.Handler) http.Handler {
|
||||
}
|
||||
}
|
||||
|
||||
// MetricsRateLimit returns middleware for /metrics that answers 429
|
||||
// Too Many Requests to a client address over the rate limit. The
|
||||
// address is the one realIP works out, so a client that is not a
|
||||
// trusted proxy cannot get a fresh allowance by sending its own
|
||||
// X-Real-IP or X-Forwarded-For. CanonicalizeIP counts all IPv6
|
||||
// addresses in one /64 as one client, since a client usually holds a
|
||||
// whole /64. An IPv4 address a proxy reports in IPv6-mapped form
|
||||
// (::ffff:203.0.113.1) is turned back into plain IPv4 first, as every
|
||||
// such address is in the same /64.
|
||||
func (m *Middleware) MetricsRateLimit() func(http.Handler) http.Handler {
|
||||
return httprate.LimitBy(
|
||||
metricsRequestLimit,
|
||||
metricsRequestWindow,
|
||||
func(request *http.Request) (string, error) {
|
||||
ip := realIP(request)
|
||||
|
||||
addr, err := netip.ParseAddr(ip)
|
||||
if err == nil {
|
||||
ip = addr.Unmap().String()
|
||||
}
|
||||
|
||||
return httprate.CanonicalizeIP(ip), nil
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
// MetricsAuth returns basic auth middleware for /metrics.
|
||||
func (m *Middleware) MetricsAuth() func(http.Handler) http.Handler {
|
||||
if m.params.Config.MetricsUsername == "" {
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"go.uber.org/fx/fxtest"
|
||||
@@ -340,3 +341,144 @@ func TestDashboardRendersWithSecurityHeaders(t *testing.T) {
|
||||
t.Errorf("CSP would block %q: %q", stylesheetPath, csp)
|
||||
}
|
||||
}
|
||||
|
||||
// Addresses for the rate limit tests: a client connecting directly, a
|
||||
// trusted proxy, and a client behind that proxy as its X-Real-IP
|
||||
// header names it.
|
||||
const (
|
||||
directClient = "198.51.100.1:4000"
|
||||
trustedProxy = "10.0.0.1:4000"
|
||||
proxiedClient = "203.0.113.1"
|
||||
)
|
||||
|
||||
// statusFrom sends a GET through handler as if from remoteAddr, with
|
||||
// an X-Real-IP header when xRealIP is not empty, and returns the
|
||||
// response status.
|
||||
func statusFrom(
|
||||
t *testing.T,
|
||||
handler http.Handler,
|
||||
remoteAddr string,
|
||||
xRealIP string,
|
||||
) int {
|
||||
t.Helper()
|
||||
|
||||
req := httptest.NewRequestWithContext(
|
||||
t.Context(), http.MethodGet, "/metrics", nil,
|
||||
)
|
||||
req.RemoteAddr = remoteAddr
|
||||
|
||||
if xRealIP != "" {
|
||||
req.Header.Set("X-Real-IP", xRealIP)
|
||||
}
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
return rec.Code
|
||||
}
|
||||
|
||||
// TestMetricsRateLimitAllowsScraping checks that one address can send,
|
||||
// within one window, what two Prometheus servers scraping every 5
|
||||
// seconds send in that time, without being turned away.
|
||||
func TestMetricsRateLimitAllowsScraping(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const scrapeInterval = 5 * time.Second
|
||||
|
||||
scrapes := 2 * int(middleware.MetricsRequestWindow/scrapeInterval)
|
||||
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
|
||||
|
||||
for i := range scrapes {
|
||||
got := statusFrom(t, limited, directClient, "")
|
||||
if got != http.StatusOK {
|
||||
t.Fatalf(
|
||||
"scrape %d of %d: status = %d, want 200",
|
||||
i+1, scrapes, got,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMetricsRateLimitKeysOnClientAddress checks which requests share
|
||||
// an allowance. Each case uses up the allowance of one client, then
|
||||
// sends one more request.
|
||||
func TestMetricsRateLimitKeysOnClientAddress(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
usedRemoteAddr string
|
||||
usedXRealIP string
|
||||
nextRemoteAddr string
|
||||
nextXRealIP string
|
||||
want int
|
||||
}{
|
||||
{
|
||||
"same address",
|
||||
directClient, "",
|
||||
directClient, "",
|
||||
http.StatusTooManyRequests,
|
||||
},
|
||||
{
|
||||
"another address",
|
||||
directClient, "",
|
||||
"198.51.100.2:4000", "",
|
||||
http.StatusOK,
|
||||
},
|
||||
{
|
||||
"own X-Real-IP from an untrusted address",
|
||||
directClient, "",
|
||||
directClient, "203.0.113.9",
|
||||
http.StatusTooManyRequests,
|
||||
},
|
||||
{
|
||||
"same client behind the proxy",
|
||||
trustedProxy, proxiedClient,
|
||||
trustedProxy, proxiedClient,
|
||||
http.StatusTooManyRequests,
|
||||
},
|
||||
{
|
||||
"another client behind the proxy",
|
||||
trustedProxy, proxiedClient,
|
||||
trustedProxy, "203.0.113.2",
|
||||
http.StatusOK,
|
||||
},
|
||||
{
|
||||
"another client behind the proxy, IPv6-mapped",
|
||||
trustedProxy, "::ffff:203.0.113.1",
|
||||
trustedProxy, "::ffff:203.0.113.2",
|
||||
http.StatusOK,
|
||||
},
|
||||
{
|
||||
"same IPv6 /64",
|
||||
"[2001:db8::1]:4000", "",
|
||||
"[2001:db8::2]:4000", "",
|
||||
http.StatusTooManyRequests,
|
||||
},
|
||||
{
|
||||
"another IPv6 /64",
|
||||
"[2001:db8::1]:4000", "",
|
||||
"[2001:db8:0:1::1]:4000", "",
|
||||
http.StatusOK,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
|
||||
|
||||
for range middleware.MetricsRequestLimit {
|
||||
statusFrom(t, limited, tt.usedRemoteAddr, tt.usedXRealIP)
|
||||
}
|
||||
|
||||
got := statusFrom(
|
||||
t, limited, tt.nextRemoteAddr, tt.nextXRealIP,
|
||||
)
|
||||
if got != tt.want {
|
||||
t.Errorf("status = %d, want %d", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -64,9 +64,12 @@ func (s *Server) SetupRoutes() {
|
||||
// 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.
|
||||
// router and its CORS. The rate limit comes before Basic Auth, so
|
||||
// failed logins count against it and a request over the limit
|
||||
// never reaches the password check.
|
||||
if s.params.Config.MetricsUsername != "" {
|
||||
metrics := chi.NewRouter()
|
||||
metrics.Use(s.mw.MetricsRateLimit())
|
||||
metrics.Use(s.mw.MetricsAuth())
|
||||
metrics.Get("/", promhttp.Handler().ServeHTTP)
|
||||
s.router.Mount("/metrics", metrics)
|
||||
|
||||
@@ -219,3 +219,72 @@ func TestPreflightAllowsOnlyWhatPublicRoutesServe(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user