package middleware import ( "net/http" "net/http/httptest" "testing" "github.com/go-chi/chi/v5" "github.com/prometheus/client_golang/prometheus/testutil" ) func TestMetricsMiddleware_RecordsRequestWithRouteTemplate(t *testing.T) { r := chi.NewRouter() r.Use(Metrics) r.Get("/things/{id}", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("ok")) }) // 用路由模板(而非具体路径)作为标签,避免高基数 before := testutil.ToFloat64(httpRequestsTotal.WithLabelValues("GET", "/things/{id}", "200")) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest("GET", "/things/42", nil)) if rec.Code != http.StatusOK { t.Fatalf("状态码应为 200,实际 %d", rec.Code) } after := testutil.ToFloat64(httpRequestsTotal.WithLabelValues("GET", "/things/{id}", "200")) if after != before+1 { t.Fatalf("请求计数应 +1:before=%v after=%v", before, after) } } func TestMetricsMiddleware_RecordsErrorStatus(t *testing.T) { r := chi.NewRouter() r.Use(Metrics) r.Get("/boom", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusInternalServerError) }) before := testutil.ToFloat64(httpRequestsTotal.WithLabelValues("GET", "/boom", "500")) rec := httptest.NewRecorder() r.ServeHTTP(rec, httptest.NewRequest("GET", "/boom", nil)) after := testutil.ToFloat64(httpRequestsTotal.WithLabelValues("GET", "/boom", "500")) if after != before+1 { t.Fatalf("500 计数应 +1:before=%v after=%v", before, after) } }