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
240
241
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
247
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
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
269 var _ Reaper = (*PromReaper)(nil)
270
271
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
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
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
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
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
375 synctest.Wait()
376 assert.ElementsMatch(t, []string{fastEnv, slowEnv}, drain(), "tick 1: both families emitted")
377
378
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
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
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
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
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