diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 3e2aeac..d34aaec 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -27,5 +27,5 @@ jobs: - name: golangci-lint uses: golangci/golangci-lint-action@1e7e51e771db61008b38414a730f564565cf7c20 # v9.2.0 with: - version: latest + version: v2.11.4 args: --timeout=10m diff --git a/go.mod b/go.mod index 7ccd310..a6667c6 100644 --- a/go.mod +++ b/go.mod @@ -20,6 +20,7 @@ require ( github.com/davecgh/go-spew v1.1.1 // indirect github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/kr/text v0.2.0 // indirect + github.com/kylelemons/godebug v1.1.0 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/prometheus/client_model v0.6.2 // indirect diff --git a/internal/middleware/metrics.go b/internal/middleware/metrics.go index 2e06fc0..e91850a 100644 --- a/internal/middleware/metrics.go +++ b/internal/middleware/metrics.go @@ -9,6 +9,8 @@ import ( "github.com/prometheus/client_golang/prometheus/promauto" ) +const unmatchedRoute = "unmatched" + var ( httpRequestsTotal = prometheus.NewCounterVec( prometheus.CounterOpts{ @@ -95,6 +97,14 @@ func (mrw *metricsResponseWriter) Write(b []byte) (int, error) { return n, err } +func routePattern(r *http.Request) string { + if r.Pattern != "" { + return r.Pattern + } + + return unmatchedRoute +} + // Metrics returns middleware that collects Prometheus metrics. func Metrics() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { @@ -108,31 +118,31 @@ func Metrics() func(http.Handler) http.Handler { bytesWritten: 0, } - // Record request size - if r.ContentLength > 0 { - httpRequestSize.WithLabelValues(r.Method, r.URL.Path).Observe(float64(r.ContentLength)) - } + contentLength := r.ContentLength - // Call next handler next.ServeHTTP(mrw, r) - // Record metrics duration := time.Since(start) + route := routePattern(r) + + if contentLength > 0 { + httpRequestSize.WithLabelValues(r.Method, route).Observe(float64(contentLength)) + } httpRequestsTotal.WithLabelValues( r.Method, - r.URL.Path, + route, strconv.Itoa(mrw.statusCode), ).Inc() httpRequestDuration.WithLabelValues( r.Method, - r.URL.Path, + route, ).Observe(duration.Seconds()) httpResponseSize.WithLabelValues( r.Method, - r.URL.Path, + route, ).Observe(float64(mrw.bytesWritten)) }) } diff --git a/internal/middleware/metrics_test.go b/internal/middleware/metrics_test.go new file mode 100644 index 0000000..eb403ca --- /dev/null +++ b/internal/middleware/metrics_test.go @@ -0,0 +1,178 @@ +package middleware + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func resetMetrics() { + httpRequestsTotal.Reset() + httpRequestDuration.Reset() + httpRequestSize.Reset() + httpResponseSize.Reset() +} + +func TestMetricsMiddleware_UsesRoutePatternNotRawURL(t *testing.T) { + resetMetrics() + + mux := http.NewServeMux() + mux.HandleFunc("GET /api/v1/{network}/bounds", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + handler := Metrics()(mux) + + for _, network := range []string{"mainnet", "sepolia", "holesky"} { + req := httptest.NewRequest(http.MethodGet, "/api/v1/"+network+"/bounds", http.NoBody) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + } + + expectedRoute := "GET /api/v1/{network}/bounds" + count := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, expectedRoute, "200")) + assert.Equal(t, float64(3), count, "all three requests should collapse to one route-pattern series") + + for _, network := range []string{"mainnet", "sepolia", "holesky"} { + rawURL := "/api/v1/" + network + "/bounds" + leak := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, rawURL, "200")) + assert.Equal(t, float64(0), leak, "raw URL %q must never become a metric label", rawURL) + } +} + +func TestMetricsMiddleware_UnmatchedRouteUsesSentinel(t *testing.T) { + resetMetrics() + + mux := http.NewServeMux() + handler := Metrics()(mux) + + req := httptest.NewRequest(http.MethodGet, "/this/path/does/not/exist", http.NoBody) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + require.Equal(t, http.StatusNotFound, rec.Code) + + count := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, unmatchedRoute, "404")) + assert.Equal(t, float64(1), count, "unmatched routes should fall back to the sentinel") + + leak := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, "/this/path/does/not/exist", "404")) + assert.Equal(t, float64(0), leak, "raw URL must never become a metric label") +} + +func TestMetricsMiddleware_CatchAllPatternCollapsesGarbage(t *testing.T) { + resetMetrics() + + mux := http.NewServeMux() + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + handler := Metrics()(mux) + + garbagePaths := []string{ + "/.env", + "/wp-login.php", + "/$$whyalwaysme@@.php", + "/%2f%2eAwS%2fCrEdEnTiAlS", + } + for _, p := range garbagePaths { + req := httptest.NewRequest(http.MethodGet, p, http.NoBody) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + } + + count := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, "/", "200")) + assert.Equal(t, float64(len(garbagePaths)), count, "all garbage URLs should collapse to the catch-all pattern") + + for _, p := range garbagePaths { + leak := testutil.ToFloat64(httpRequestsTotal.WithLabelValues(http.MethodGet, p, "200")) + assert.Equal(t, float64(0), leak, "scanner garbage URL %q must never become a metric label", p) + } +} + +func TestMetricsMiddleware_RequestSizeRecordedAgainstPattern(t *testing.T) { + resetMetrics() + + mux := http.NewServeMux() + mux.HandleFunc("POST /api/v1/upload", func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + + w.WriteHeader(http.StatusAccepted) + }) + + handler := Metrics()(mux) + + body := strings.Repeat("x", 1024) + req := httptest.NewRequest(http.MethodPost, "/api/v1/upload", strings.NewReader(body)) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + require.Equal(t, http.StatusAccepted, rec.Code) + + count := testutil.CollectAndCount(httpRequestSize, "http_request_size_bytes") + assert.Equal(t, 1, count, "request size should be observed exactly once against the route pattern") +} + +func TestMetricsMiddleware_MultipleRoutesGetDistinctLabels(t *testing.T) { + resetMetrics() + + mux := http.NewServeMux() + mux.HandleFunc("GET /health", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + mux.HandleFunc("GET /api/v1/config", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + mux.HandleFunc("GET /api/v1/{network}/bounds", func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + handler := Metrics()(mux) + + for _, p := range []string{"/health", "/api/v1/config", "/api/v1/mainnet/bounds", "/api/v1/sepolia/bounds"} { + req := httptest.NewRequest(http.MethodGet, p, http.NoBody) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + } + + assert.Equal(t, float64(1), testutil.ToFloat64( + httpRequestsTotal.WithLabelValues(http.MethodGet, "GET /health", "200"), + )) + assert.Equal(t, float64(1), testutil.ToFloat64( + httpRequestsTotal.WithLabelValues(http.MethodGet, "GET /api/v1/config", "200"), + )) + assert.Equal(t, float64(2), testutil.ToFloat64( + httpRequestsTotal.WithLabelValues(http.MethodGet, "GET /api/v1/{network}/bounds", "200"), + )) +} + +func TestRoutePattern_ReturnsSentinelWhenUnset(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/anything", http.NoBody) + assert.Equal(t, unmatchedRoute, routePattern(req)) +} + +func TestRoutePattern_ReturnsMatchedPattern(t *testing.T) { + mux := http.NewServeMux() + + var observed string + + mux.HandleFunc("GET /api/v1/{network}/bounds", func(w http.ResponseWriter, r *http.Request) { + observed = routePattern(r) + + w.WriteHeader(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/mainnet/bounds", http.NoBody) + rec := httptest.NewRecorder() + mux.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + assert.Equal(t, "GET /api/v1/{network}/bounds", observed) +}