...

Source file src/github.com/cybertec-postgresql/pgwatch/v6/internal/reaper/prometheus_test.go

Documentation: github.com/cybertec-postgresql/pgwatch/v6/internal/reaper

     1  package reaper
     2  
     3  import (
     4  	"io"
     5  	"math"
     6  	"net/http"
     7  	"net/http/httptest"
     8  	"strings"
     9  	"sync/atomic"
    10  	"testing"
    11  	"testing/synctest"
    12  	"time"
    13  
    14  	"github.com/cybertec-postgresql/pgwatch/v6/internal/log"
    15  	"github.com/cybertec-postgresql/pgwatch/v6/internal/metrics"
    16  	"github.com/cybertec-postgresql/pgwatch/v6/internal/sources"
    17  	"github.com/cybertec-postgresql/pgwatch/v6/internal/testutil"
    18  	"github.com/sirupsen/logrus"
    19  	"github.com/stretchr/testify/assert"
    20  	"github.com/stretchr/testify/require"
    21  )
    22  
    23  const scrapeAllFixture = `# HELP go_goroutines Number of goroutines
    24  # TYPE go_goroutines gauge
    25  go_goroutines{job="pgwatch",instance="local"} 42 1700000000000
    26  # HELP http_requests_total Total HTTP requests
    27  # TYPE http_requests_total counter
    28  http_requests_total{method="GET",code="200"} 1024
    29  `
    30  
    31  func TestScrapeAll(t *testing.T) {
    32  	testCases := []struct {
    33  		name  string
    34  		check func(t *testing.T, got []metrics.MeasurementEnvelope)
    35  	}{
    36  		{
    37  			name: "returns expected measurement count",
    38  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    39  				t.Helper()
    40  				require.Len(t, got, 2)
    41  				names := make([]string, len(got))
    42  				for i, e := range got {
    43  					names[i] = e.MetricName
    44  				}
    45  				assert.ElementsMatch(t, []string{"go_goroutines", "http_requests_total"}, names)
    46  				require.Len(t, requireEnvelope(t, got, "go_goroutines").Data, 1)
    47  				require.Len(t, requireEnvelope(t, got, "http_requests_total").Data, 1)
    48  			},
    49  		},
    50  		{
    51  			name: "source kind set",
    52  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    53  				t.Helper()
    54  				assert.Equal(t, string(sources.SourcePrometheus), requireEnvelope(t, got, "go_goroutines").SourceKind)
    55  			},
    56  		},
    57  		{
    58  			name: "tag labels present",
    59  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    60  				t.Helper()
    61  				measurement := requireSingleMeasurement(t, requireEnvelope(t, got, "go_goroutines").Data)
    62  				assert.Equal(t, "pgwatch", measurement[metrics.TagPrefix+"job"])
    63  				assert.Equal(t, "local", measurement[metrics.TagPrefix+"instance"])
    64  			},
    65  		},
    66  		{
    67  			name: "no __name__ label",
    68  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    69  				t.Helper()
    70  				measurement := requireSingleMeasurement(t, requireEnvelope(t, got, "go_goroutines").Data)
    71  				assert.NotContains(t, measurement, metrics.TagPrefix+"__name__")
    72  				assert.NotContains(t, measurement, "__name__")
    73  			},
    74  		},
    75  		{
    76  			name: "value column",
    77  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    78  				t.Helper()
    79  				measurement := requireSingleMeasurement(t, requireEnvelope(t, got, "go_goroutines").Data)
    80  				assert.Equal(t, float64(42), measurement["go_goroutines"])
    81  			},
    82  		},
    83  		{
    84  			name: "epoch_ns from timestamp",
    85  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    86  				t.Helper()
    87  				measurement := requireSingleMeasurement(t, requireEnvelope(t, got, "go_goroutines").Data)
    88  				assert.Equal(t, int64(1700000000000*1_000_000), measurement[metrics.EpochColumnName])
    89  			},
    90  		},
    91  		{
    92  			name: "epoch_ns fallback",
    93  			check: func(t *testing.T, got []metrics.MeasurementEnvelope) {
    94  				t.Helper()
    95  				measurement := requireSingleMeasurement(t, requireEnvelope(t, got, "http_requests_total").Data)
    96  				epoch, ok := measurement[metrics.EpochColumnName].(int64)
    97  				require.True(t, ok)
    98  				assert.WithinDuration(t, time.Now(), time.Unix(0, epoch), time.Second)
    99  			},
   100  		},
   101  	}
   102  
   103  	for _, tc := range testCases {
   104  		t.Run(tc.name, func(t *testing.T) {
   105  			sc := newTestPromConn(t, scrapeAllFixture)
   106  			got, err := (&PromReaper{md: sc}).ScrapeAll(t.Context())
   107  			require.NoError(t, err)
   108  			tc.check(t, got)
   109  		})
   110  	}
   111  }
   112  
   113  func TestScrapeAll_NonFinite(t *testing.T) {
   114  	const fixture = `# HELP mymetric A metric
   115  # TYPE mymetric gauge
   116  mymetric{} +Inf
   117  mymetric{foo="bar"} -Inf
   118  mymetric{baz="qux"} NaN
   119  `
   120  
   121  	sc := newTestPromConn(t, fixture)
   122  	got, err := (&PromReaper{md: sc}).ScrapeAll(t.Context())
   123  	require.NoError(t, err)
   124  
   125  	env := requireEnvelope(t, got, "mymetric")
   126  	require.Len(t, env.Data, 3)
   127  
   128  	testCases := []struct {
   129  		name  string
   130  		tags  map[string]string
   131  		check func(float64) bool
   132  	}{
   133  		{
   134  			name:  "positive infinity",
   135  			tags:  map[string]string{},
   136  			check: func(v float64) bool { return math.IsInf(v, 1) },
   137  		},
   138  		{
   139  			name:  "negative infinity",
   140  			tags:  map[string]string{"foo": "bar"},
   141  			check: func(v float64) bool { return math.IsInf(v, -1) },
   142  		},
   143  		{
   144  			name:  "nan",
   145  			tags:  map[string]string{"baz": "qux"},
   146  			check: math.IsNaN,
   147  		},
   148  	}
   149  
   150  	for _, tc := range testCases {
   151  		t.Run(tc.name, func(t *testing.T) {
   152  			measurement := requireMeasurementByTags(t, env.Data, tc.tags)
   153  			value, ok := measurement["mymetric"].(float64)
   154  			require.True(t, ok)
   155  			assert.True(t, tc.check(value))
   156  		})
   157  	}
   158  }
   159  
   160  func TestScrapeAll_AcceptHeader(t *testing.T) {
   161  	var accept string
   162  
   163  	srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   164  		accept = r.Header.Get("Accept")
   165  		w.Header().Set("Content-Type", "text/plain; version=0.0.4")
   166  		_, _ = w.Write([]byte("# HELP up Up\n# TYPE up gauge\nup 1\n"))
   167  	}))
   168  	t.Cleanup(srv.Close)
   169  
   170  	sc := &sources.PromConn{
   171  		Source:     sources.Source{ConnStr: srv.URL},
   172  		HTTPClient: srv.Client(),
   173  	}
   174  	require.NoError(t, sc.ParseConfig())
   175  
   176  	_, err := (&PromReaper{md: sc}).ScrapeAll(t.Context())
   177  	require.NoError(t, err)
   178  	assert.Equal(t, "text/plain", accept)
   179  }
   180  
   181  func newTestPromConn(t *testing.T, body string) *sources.PromConn {
   182  	t.Helper()
   183  	srv := testutil.NewFakeExporter(t, body)
   184  	pc := &sources.PromConn{
   185  		Source: sources.Source{
   186  			Kind:    sources.SourcePrometheus,
   187  			ConnStr: srv.URL,
   188  		},
   189  		HTTPClient: srv.Client(),
   190  	}
   191  	require.NoError(t, pc.ParseConfig())
   192  	return pc
   193  }
   194  
   195  func requireEnvelope(t *testing.T, envelopes []metrics.MeasurementEnvelope, metricName string) metrics.MeasurementEnvelope {
   196  	t.Helper()
   197  	for _, e := range envelopes {
   198  		if e.MetricName == metricName {
   199  			return e
   200  		}
   201  	}
   202  	require.FailNowf(t, "envelope not found", "no envelope with MetricName %q", metricName)
   203  	return metrics.MeasurementEnvelope{}
   204  }
   205  
   206  func requireSingleMeasurement(t *testing.T, data metrics.Measurements) metrics.Measurement {
   207  	t.Helper()
   208  	require.Len(t, data, 1)
   209  	return metrics.Measurement(data[0])
   210  }
   211  
   212  func requireMeasurementByTags(t *testing.T, measurements metrics.Measurements, tags map[string]string) metrics.Measurement {
   213  	t.Helper()
   214  	for _, measurement := range measurements {
   215  		if hasExactTags(measurement, tags) {
   216  			return metrics.Measurement(measurement)
   217  		}
   218  	}
   219  	require.FailNowf(t, "measurement not found", "expected tags %v", tags)
   220  	return nil
   221  }
   222  
   223  func hasExactTags(measurement map[string]any, tags map[string]string) bool {
   224  	tagCount := 0
   225  	for key, value := range measurement {
   226  		label, ok := strings.CutPrefix(key, metrics.TagPrefix)
   227  		if !ok {
   228  			continue
   229  		}
   230  		tagCount++
   231  		expected, ok := tags[label]
   232  		if !ok || value != expected {
   233  			return false
   234  		}
   235  	}
   236  	return tagCount == len(tags)
   237  }
   238  
   239  // roundTripFunc is a fake http.RoundTripper backed by a plain function.
   240  // Using it instead of a real httptest.Server avoids spawning IO-blocked goroutines
   241  // inside a synctest bubble (which would prevent synctest.Wait from ever returning).
   242  type roundTripFunc func(*http.Request) (*http.Response, error)
   243  
   244  func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
   245  
   246  // fakePromConn builds a PromConn whose HTTP client uses fn as its transport.
   247  // ConnStr can be any syntactically valid URL; fn intercepts every request.
   248  func fakePromConn(t *testing.T, src sources.Source, fn roundTripFunc) *sources.PromConn {
   249  	t.Helper()
   250  	src.ConnStr = "http://fake.local"
   251  	pc := &sources.PromConn{
   252  		Source:     src,
   253  		HTTPClient: &http.Client{Transport: fn},
   254  	}
   255  	require.NoError(t, pc.ParseConfig())
   256  	return pc
   257  }
   258  
   259  // fakeResponse builds a minimal 200 Prometheus text-format response.
   260  func fakeResponse(body string) *http.Response {
   261  	return &http.Response{
   262  		StatusCode: 200,
   263  		Header:     http.Header{"Content-Type": []string{"text/plain; version=0.0.4"}},
   264  		Body:       io.NopCloser(strings.NewReader(body)),
   265  	}
   266  }
   267  
   268  // compile-time check that PromReaper implements Reaper.
   269  var _ Reaper = (*PromReaper)(nil)
   270  
   271  // TestCalcScrapeInterval verifies GCD-based interval calculation and defaults.
   272  func TestCalcScrapeInterval(t *testing.T) {
   273  	tests := []struct {
   274  		name    string
   275  		metrics metrics.MetricIntervals
   276  		want    time.Duration
   277  	}{
   278  		{
   279  			name:    "multiple intervals produce GCD",
   280  			metrics: metrics.MetricIntervals{"a": 30, "b": 60},
   281  			want:    30 * time.Second,
   282  		},
   283  		{
   284  			name:    "single interval returns that value",
   285  			metrics: metrics.MetricIntervals{"a": 60},
   286  			want:    60 * time.Second,
   287  		},
   288  		{
   289  			name:    "empty intervals scrape-all defaults to 60s",
   290  			metrics: metrics.MetricIntervals{},
   291  			want:    defaultScrapeInterval,
   292  		},
   293  		{
   294  			name:    "all intervals below min floored to minTickInterval",
   295  			metrics: metrics.MetricIntervals{"a": 0, "b": -1},
   296  			want:    minTickInterval * time.Second,
   297  		},
   298  	}
   299  
   300  	for _, tc := range tests {
   301  		t.Run(tc.name, func(t *testing.T) {
   302  			pr := &PromReaper{
   303  				md: &sources.PromConn{
   304  					Source: sources.Source{Metrics: tc.metrics},
   305  				},
   306  			}
   307  			assert.Equal(t, tc.want, pr.calcScrapeInterval())
   308  		})
   309  	}
   310  }
   311  
   312  // when Metrics is empty, Reap uses 60 s interval and logs a warning.
   313  func TestPromReaper_ScrapeAllMode(t *testing.T) {
   314  	synctest.Test(t, func(t *testing.T) {
   315  		var scrapeCount atomic.Int32
   316  
   317  		ctx, out := testutil.NewTestLogger(t, logrus.WarnLevel)
   318  
   319  		md := fakePromConn(t, sources.Source{Name: "prom_scrapeall", Kind: sources.SourcePrometheus}, func(*http.Request) (*http.Response, error) {
   320  			scrapeCount.Add(1)
   321  			return fakeResponse("# HELP up Up\n# TYPE up gauge\nup 1\n"), nil
   322  		})
   323  
   324  		pr := NewPromSourceReaper(&reaper{measurementCh: make(chan metrics.MeasurementEnvelope, 100)}, md)
   325  		go pr.Reap(ctx)
   326  
   327  		// t=0: first scrape completes; goroutine blocks on time.After(60s).
   328  		synctest.Wait()
   329  		assert.Equal(t, int32(1), scrapeCount.Load(), "first scrape should occur at t=0")
   330  		assert.Contains(t, out.String(), "scrape-all", "scrape-all mode warning should be logged")
   331  
   332  		// Advance 60 s → second scrape.
   333  		time.Sleep(60 * time.Second)
   334  		synctest.Wait()
   335  		assert.Equal(t, int32(2), scrapeCount.Load(), "second scrape should occur after 60 s")
   336  	})
   337  }
   338  
   339  // per-family emit gating — families respect their configured intervals.
   340  func TestPromReaper_EmitGating(t *testing.T) {
   341  	const gatingBody = "# HELP fast Fast metric\n# TYPE fast gauge\nfast 1\n" +
   342  		"# HELP slow Slow metric\n# TYPE slow gauge\nslow 2\n"
   343  
   344  	synctest.Test(t, func(t *testing.T) {
   345  		const fastEnv = "fast"
   346  		const slowEnv = "slow"
   347  
   348  		r := &reaper{measurementCh: make(chan metrics.MeasurementEnvelope, 100)}
   349  		md := fakePromConn(t, sources.Source{
   350  			Name:    "gating_test",
   351  			Kind:    sources.SourcePrometheus,
   352  			Metrics: metrics.MetricIntervals{"fast": 30, "slow": 60},
   353  		}, func(*http.Request) (*http.Response, error) {
   354  			return fakeResponse(gatingBody), nil
   355  		})
   356  
   357  		pr := NewPromSourceReaper(r, md)
   358  
   359  		ctx := log.WithLogger(t.Context(), log.NewNoopLogger())
   360  		go pr.Reap(ctx)
   361  
   362  		drain := func() []string {
   363  			var names []string
   364  			for {
   365  				select {
   366  				case env := <-r.measurementCh:
   367  					names = append(names, env.MetricName)
   368  				default:
   369  					return names
   370  				}
   371  			}
   372  		}
   373  
   374  		// Tick 1 (t=0): lastEmitted is zero for both → both emitted.
   375  		synctest.Wait()
   376  		assert.ElementsMatch(t, []string{fastEnv, slowEnv}, drain(), "tick 1: both families emitted")
   377  
   378  		// Tick 2 (t=30s): only "fast" due; "slow" still within its 60 s window.
   379  		time.Sleep(30*time.Second + time.Millisecond)
   380  		synctest.Wait()
   381  		assert.ElementsMatch(t, []string{fastEnv}, drain(), "tick 2: only 'fast' emitted")
   382  
   383  		// Tick 3 (t=60s): both families now due.
   384  		time.Sleep(30 * time.Second)
   385  		synctest.Wait()
   386  		assert.ElementsMatch(t, []string{fastEnv, slowEnv}, drain(), "tick 3: both families emitted again")
   387  	})
   388  }
   389  
   390  // ScrapeAll error path — warning is logged, lastEmitted stays unchanged, loop continues.
   391  func TestPromReaper_ScrapeError(t *testing.T) {
   392  	synctest.Test(t, func(t *testing.T) {
   393  		ctx, out := testutil.NewTestLogger(t, logrus.WarnLevel)
   394  		r := &reaper{measurementCh: make(chan metrics.MeasurementEnvelope, 10)}
   395  
   396  		md := fakePromConn(t, sources.Source{
   397  			Name:    "error_test",
   398  			Kind:    sources.SourcePrometheus,
   399  			Metrics: metrics.MetricIntervals{"some_metric": 30},
   400  		}, func(*http.Request) (*http.Response, error) {
   401  			return &http.Response{
   402  				StatusCode: http.StatusInternalServerError,
   403  				Status:     "500 Internal Server Error",
   404  				Body:       io.NopCloser(strings.NewReader("")),
   405  			}, nil
   406  		})
   407  
   408  		pr := NewPromSourceReaper(r, md)
   409  		go pr.Reap(ctx)
   410  
   411  		// First (failing) scrape completes; goroutine blocks on time.After.
   412  		synctest.Wait()
   413  		assert.Contains(t, out.String(), "scrape failed", "warning should be logged on scrape error")
   414  		assert.Empty(t, pr.lastEmitted, "lastEmitted should not change on error")
   415  
   416  		select {
   417  		case <-r.measurementCh:
   418  			t.Error("no measurements should be dispatched on scrape error")
   419  		default:
   420  		}
   421  
   422  		// Advance to trigger a second tick; loop must continue without crashing.
   423  		time.Sleep(30*time.Second + time.Millisecond)
   424  		synctest.Wait()
   425  
   426  		select {
   427  		case <-r.measurementCh:
   428  			t.Error("no measurements should be dispatched on repeated scrape errors")
   429  		default:
   430  		}
   431  	})
   432  }
   433