Skip to main content

measured_prometheus_protobuf/
lib.rs

1#![allow(clippy::cast_precision_loss)]
2
3use std::io::Write;
4
5use encoding::{encode_key, encode_varint, encoded_len_varint, key_len, WireType::LengthDelimited};
6use measured::{
7    label::{LabelGroupVisitor, LabelName, LabelValue, LabelVisitor},
8    metric::{
9        counter::CounterState,
10        gauge::{FloatGaugeState, GaugeState},
11        group::Encoding,
12        name::MetricNameEncoder,
13        MetricEncoding,
14    },
15    LabelGroup,
16};
17
18mod encoding;
19
20/// The prometheus text encoder helper
21pub struct ProtoEncoder<W> {
22    state: State,
23    pub writer: W,
24    buf: Vec<u8>,
25}
26
27impl<W: Write> ProtoEncoder<W> {
28    /// Create a new text encoder.
29    ///
30    /// This should ideally be cached and re-used between collections to reduce re-allocating
31    pub fn new(w: W) -> Self {
32        Self {
33            state: State::Init,
34            writer: w,
35            buf: Vec::new(),
36        }
37    }
38
39    /// Finish the text encoding and extract the bytes to send in a HTTP response.
40    pub fn flush(&mut self) -> std::io::Result<()> {
41        self.flush_buf()?;
42        self.writer.flush()
43    }
44
45    fn flush_buf(&mut self) -> Result<(), std::io::Error> {
46        if self.state == State::Metrics {
47            self.state = State::Init;
48
49            let len = self.buf.len() - 10;
50            let varint_len = encoded_len_varint(len as u64);
51            let offset = 10 - varint_len;
52            encode_varint(len as u64, &mut &mut self.buf[offset..]);
53            self.writer.write_all(&self.buf[offset..])?;
54
55            self.buf.resize(10, 0);
56        } else if self.buf.is_empty() {
57            self.buf.resize(10, 0);
58        }
59        debug_assert!(self.buf.len() >= 10);
60        Ok(())
61    }
62}
63
64#[derive(Clone, Copy, Debug, PartialEq)]
65enum State {
66    Init,
67    Help,
68    Metrics,
69}
70
71/// Prometheus only supports these 5 types of metrics
72#[derive(Clone, Copy, Debug)]
73pub enum MetricType {
74    /// Corresponds to [`Counter`](crate::Counter)
75    Counter,
76    /// Corresponds to [`Histogram`](crate::Histogram)
77    Histogram,
78    /// Corresponds to [`Gauge`](crate::Gauge)
79    Gauge,
80    /// Not currently supported
81    Summary,
82    /// Not currently supported
83    Untyped,
84}
85
86impl<W: Write> Encoding for ProtoEncoder<W> {
87    type Err = std::io::Error;
88
89    const MIME_TYPE: &'static str = "application/vnd.google.protobuf; proto=io.prometheus.client.MetricFamily; encoding=delimited";
90
91    /// Write the help line for a metric
92    fn write_help(
93        &mut self,
94        name: impl MetricNameEncoder,
95        help: &str,
96    ) -> Result<(), std::io::Error> {
97        self.flush_buf()?;
98
99        // optional string     name   = 1;
100        encode_key(1, LengthDelimited, &mut self.buf);
101        encode_varint(name.encode_len() as u64, &mut self.buf);
102        name.encode_utf8(&mut self.buf)?;
103
104        // optional string     help   = 2;
105        encoding::encode_str(2, help, &mut self.buf);
106
107        self.state = State::Help;
108
109        Ok(())
110    }
111}
112
113struct LenVisitor {}
114impl LabelVisitor for LenVisitor {
115    type Output = usize;
116    fn write_int(self, x: i64) -> usize {
117        encoding::encoded_len_str(2, itoa::Buffer::new().format(x))
118    }
119
120    fn write_float(self, x: f64) -> usize {
121        if x.is_infinite() {
122            if x.is_sign_positive() {
123                encoding::encoded_len_str(2, "+Inf")
124            } else {
125                encoding::encoded_len_str(2, "-Inf")
126            }
127        } else if x.is_nan() {
128            encoding::encoded_len_str(2, "NaN")
129        } else {
130            encoding::encoded_len_str(2, ryu::Buffer::new().format(x))
131        }
132    }
133
134    fn write_str(self, x: &str) -> usize {
135        encoding::encoded_len_str(2, x)
136    }
137}
138
139struct GroupLenVisitor {
140    len: usize,
141}
142impl LabelGroupVisitor for GroupLenVisitor {
143    type Output = ();
144    fn write_value(&mut self, name: &LabelName, x: &impl LabelValue) {
145        let mut label_pair_len = 0;
146
147        label_pair_len += encoding::encoded_len_str(1, name.as_str());
148        label_pair_len += x.visit(LenVisitor {});
149
150        self.len += message_len(1, label_pair_len);
151    }
152}
153
154struct Visitor<'a> {
155    buf: &'a mut Vec<u8>,
156}
157impl LabelVisitor for Visitor<'_> {
158    type Output = ();
159    fn write_int(self, x: i64) {
160        self.write_str(itoa::Buffer::new().format(x));
161    }
162
163    fn write_float(self, x: f64) {
164        if x.is_infinite() {
165            if x.is_sign_positive() {
166                self.write_str("+Inf");
167            } else {
168                self.write_str("-Inf");
169            }
170        } else if x.is_nan() {
171            self.write_str("NaN");
172        } else {
173            self.write_str(ryu::Buffer::new().format(x));
174        }
175    }
176
177    fn write_str(self, x: &str) {
178        // optional string value = 2;
179        encoding::encode_str(2, x, self.buf);
180    }
181}
182
183fn encode_message(
184    tag: u32,
185    len: usize,
186    buf: &mut Vec<u8>,
187    message: impl for<'a> FnOnce(&'a mut Vec<u8>),
188) {
189    encode_key(tag, LengthDelimited, buf);
190    encode_varint(len as u64, buf);
191
192    {
193        let offset = buf.len();
194        message(buf);
195        debug_assert_eq!(buf.len() - offset, len);
196    }
197}
198
199fn message_len(tag: u32, len: usize) -> usize {
200    key_len(tag) + encoded_len_varint(len as u64) + len
201}
202
203struct GroupVisitor<'a> {
204    buf: &'a mut Vec<u8>,
205}
206impl LabelGroupVisitor for GroupVisitor<'_> {
207    type Output = ();
208    fn write_value(&mut self, name: &LabelName, x: &impl LabelValue) {
209        let mut label_pair_len = 0;
210        label_pair_len += encoding::encoded_len_str(1, name.as_str());
211        label_pair_len += x.visit(LenVisitor {});
212
213        // repeated LabelPair label        = 1;
214        encode_message(1, label_pair_len, self.buf, |buf| {
215            // optional string name  = 1;
216            encoding::encode_str(1, name.as_str(), buf);
217
218            x.visit(Visitor { buf });
219        });
220    }
221}
222
223impl<W: Write> MetricEncoding<ProtoEncoder<W>> for CounterState {
224    fn write_type(
225        name: impl MetricNameEncoder,
226        enc: &mut ProtoEncoder<W>,
227    ) -> Result<(), std::io::Error> {
228        enc.flush_buf()?;
229
230        if enc.state == State::Init {
231            // optional string     name   = 1;
232            encode_key(1, LengthDelimited, &mut enc.buf);
233            encode_varint(name.encode_len() as u64, &mut enc.buf);
234            name.encode_utf8(&mut enc.buf)?;
235        }
236
237        // optional MetricType type   = 3;
238        // COUNTER = 0;
239        encoding::encode_i32(3, 0, &mut enc.buf);
240
241        Ok(())
242    }
243
244    fn collect_into(
245        &self,
246        _m: &(),
247        labels: impl LabelGroup,
248        _name: impl MetricNameEncoder,
249        enc: &mut ProtoEncoder<W>,
250    ) -> Result<(), std::io::Error> {
251        enc.state = State::Metrics;
252
253        let mut metric_len = 0;
254
255        let mut label_pairs_len = GroupLenVisitor { len: 0 };
256        labels.visit_values(&mut label_pairs_len);
257        metric_len += label_pairs_len.len;
258
259        let count = self.count.load(std::sync::atomic::Ordering::Relaxed) as f64;
260        let count_len = encoding::encoded_len_f64(1, count);
261        metric_len += message_len(3, count_len);
262
263        // repeated Metric     metric = 4;
264        encode_message(4, metric_len, &mut enc.buf, |buf| {
265            labels.visit_values(&mut GroupVisitor { buf });
266
267            // optional Counter   counter      = 3;
268            encode_message(3, count_len, buf, |buf| {
269                // optional double   value    = 1;
270                encoding::encode_f64(1, count, buf);
271            });
272        });
273
274        Ok(())
275    }
276}
277
278impl<W: Write> MetricEncoding<ProtoEncoder<W>> for GaugeState {
279    fn write_type(
280        name: impl MetricNameEncoder,
281        enc: &mut ProtoEncoder<W>,
282    ) -> Result<(), std::io::Error> {
283        enc.flush_buf()?;
284
285        if enc.state == State::Init {
286            // optional string     name   = 1;
287            encode_key(1, LengthDelimited, &mut enc.buf);
288            encode_varint(name.encode_len() as u64, &mut enc.buf);
289            name.encode_utf8(&mut enc.buf)?;
290        }
291
292        // optional MetricType type   = 3;
293        // GAUGE = 1;
294        encoding::encode_i32(3, 1, &mut enc.buf);
295
296        Ok(())
297    }
298
299    fn collect_into(
300        &self,
301        _m: &(),
302        labels: impl LabelGroup,
303        _name: impl MetricNameEncoder,
304        enc: &mut ProtoEncoder<W>,
305    ) -> Result<(), std::io::Error> {
306        enc.state = State::Metrics;
307
308        let mut metric_len = 0;
309
310        let mut label_pairs_len = GroupLenVisitor { len: 0 };
311        labels.visit_values(&mut label_pairs_len);
312        metric_len += label_pairs_len.len;
313
314        let gauge = self.count.load(std::sync::atomic::Ordering::Relaxed) as f64;
315        let gauge_len = encoding::encoded_len_f64(1, gauge);
316        metric_len += message_len(3, gauge_len);
317
318        // repeated Metric     metric = 4;
319        encode_message(4, metric_len, &mut enc.buf, |buf| {
320            labels.visit_values(&mut GroupVisitor { buf });
321
322            // optional Gauge   gauge      = 2;
323            encode_message(2, gauge_len, buf, |buf| {
324                // optional double   value    = 1;
325                encoding::encode_f64(1, gauge, buf);
326            });
327        });
328
329        Ok(())
330    }
331}
332
333impl<W: Write> MetricEncoding<ProtoEncoder<W>> for FloatGaugeState {
334    fn write_type(
335        name: impl MetricNameEncoder,
336        enc: &mut ProtoEncoder<W>,
337    ) -> Result<(), std::io::Error> {
338        enc.flush_buf()?;
339
340        if enc.state == State::Init {
341            // optional string     name   = 1;
342            encode_key(1, LengthDelimited, &mut enc.buf);
343            encode_varint(name.encode_len() as u64, &mut enc.buf);
344            name.encode_utf8(&mut enc.buf)?;
345        }
346
347        // optional MetricType type   = 3;
348        // GAUGE = 1;
349        encoding::encode_i32(3, 1, &mut enc.buf);
350
351        Ok(())
352    }
353
354    fn collect_into(
355        &self,
356        _m: &(),
357        labels: impl LabelGroup,
358        _name: impl MetricNameEncoder,
359        enc: &mut ProtoEncoder<W>,
360    ) -> Result<(), std::io::Error> {
361        enc.state = State::Metrics;
362
363        let mut metric_len = 0;
364
365        let mut label_pairs_len = GroupLenVisitor { len: 0 };
366        labels.visit_values(&mut label_pairs_len);
367        metric_len += label_pairs_len.len;
368
369        let gauge = self.count.get();
370        let gauge_len = encoding::encoded_len_f64(1, gauge);
371        metric_len += message_len(3, gauge_len);
372
373        // repeated Metric     metric = 4;
374        encode_message(4, metric_len, &mut enc.buf, |buf| {
375            labels.visit_values(&mut GroupVisitor { buf });
376
377            // optional Gauge   gauge      = 2;
378            encode_message(2, gauge_len, buf, |buf| {
379                // optional double   value    = 1;
380                encoding::encode_f64(1, gauge, buf);
381            });
382        });
383
384        Ok(())
385    }
386}
387
388#[cfg(test)]
389mod generated;
390
391#[cfg(test)]
392mod tests {
393    use std::vec;
394
395    use bytes::{BufMut, BytesMut};
396    use measured::{
397        metric::{
398            group::Encoding,
399            name::{MetricName, Total},
400            MetricFamilyEncoding,
401        },
402        CounterVec, GaugeVec,
403    };
404    use prost::Message;
405
406    use crate::{
407        generated::{Counter, Gauge, LabelPair, Metric, MetricFamily, MetricType},
408        ProtoEncoder,
409    };
410
411    #[derive(Clone, Copy, PartialEq, Debug, measured::LabelGroup)]
412    #[label(set = RequestLabelSet)]
413    struct RequestLabels {
414        method: Method,
415        code: StatusCode,
416    }
417
418    #[derive(Clone, Copy, PartialEq, Debug, measured::FixedCardinalityLabel)]
419    #[label(rename_all = "snake_case")]
420    enum Method {
421        Post,
422        Get,
423    }
424
425    #[derive(Clone, Copy, PartialEq, Debug, measured::FixedCardinalityLabel)]
426    enum StatusCode {
427        Ok = 200,
428        BadRequest = 400,
429    }
430
431    #[test]
432    fn counters() {
433        let requests = CounterVec::<RequestLabelSet>::new();
434
435        let labels = RequestLabels {
436            method: Method::Post,
437            code: StatusCode::Ok,
438        };
439        requests.inc_by(labels, 1027);
440
441        let labels = RequestLabels {
442            method: Method::Get,
443            code: StatusCode::BadRequest,
444        };
445        requests.inc_by(labels, 3);
446
447        let mut enc = ProtoEncoder::new(BytesMut::new().writer());
448
449        let name = MetricName::from_str("http_request").with_suffix(Total);
450        enc.write_help(&name, "The total number of HTTP requests.")
451            .unwrap();
452        requests.collect_family_into(&name, &mut enc).unwrap();
453        enc.flush().unwrap();
454        let actual_msg = enc.writer.into_inner();
455
456        let expected = MetricFamily {
457            name: Some("http_request_total".to_string()),
458            help: Some("The total number of HTTP requests.".to_string()),
459            r#type: Some(MetricType::Counter as i32),
460            metric: vec![
461                Metric {
462                    label: vec![
463                        LabelPair {
464                            name: Some("method".to_owned()),
465                            value: Some("post".to_owned()),
466                        },
467                        LabelPair {
468                            name: Some("code".to_owned()),
469                            value: Some("200".to_owned()),
470                        },
471                    ],
472                    gauge: None,
473                    counter: Some(Counter {
474                        value: Some(1027.0),
475                        exemplar: None,
476                        created_timestamp: None,
477                    }),
478                    summary: None,
479                    untyped: None,
480                    histogram: None,
481                    timestamp_ms: None,
482                },
483                Metric {
484                    label: vec![
485                        LabelPair {
486                            name: Some("method".to_owned()),
487                            value: Some("get".to_owned()),
488                        },
489                        LabelPair {
490                            name: Some("code".to_owned()),
491                            value: Some("400".to_owned()),
492                        },
493                    ],
494                    gauge: None,
495                    counter: Some(Counter {
496                        value: Some(3.0),
497                        exemplar: None,
498                        created_timestamp: None,
499                    }),
500                    summary: None,
501                    untyped: None,
502                    histogram: None,
503                    timestamp_ms: None,
504                },
505            ],
506            unit: None,
507        };
508        let mut expected_msg = BytesMut::new();
509        expected.encode_length_delimited(&mut expected_msg).unwrap();
510
511        assert_eq!(actual_msg, expected_msg);
512
513        let actual = MetricFamily::decode_length_delimited(actual_msg).unwrap();
514        assert_eq!(actual, expected);
515    }
516
517    #[test]
518    fn gauge() {
519        let requests = GaugeVec::<RequestLabelSet>::new();
520
521        let labels = RequestLabels {
522            method: Method::Post,
523            code: StatusCode::Ok,
524        };
525        requests.inc_by(labels, 1027);
526
527        let labels = RequestLabels {
528            method: Method::Get,
529            code: StatusCode::BadRequest,
530        };
531        requests.inc_by(labels, 3);
532
533        let mut enc = ProtoEncoder::new(BytesMut::new().writer());
534
535        let name = MetricName::from_str("http_request").with_suffix(Total);
536        enc.write_help(&name, "The total number of HTTP requests.")
537            .unwrap();
538        requests.collect_family_into(&name, &mut enc).unwrap();
539        enc.flush().unwrap();
540        let actual_msg = enc.writer.into_inner();
541
542        let expected = MetricFamily {
543            name: Some("http_request_total".to_string()),
544            help: Some("The total number of HTTP requests.".to_string()),
545            r#type: Some(MetricType::Gauge as i32),
546            metric: vec![
547                Metric {
548                    label: vec![
549                        LabelPair {
550                            name: Some("method".to_owned()),
551                            value: Some("post".to_owned()),
552                        },
553                        LabelPair {
554                            name: Some("code".to_owned()),
555                            value: Some("200".to_owned()),
556                        },
557                    ],
558                    gauge: Some(Gauge {
559                        value: Some(1027.0),
560                    }),
561                    counter: None,
562                    summary: None,
563                    untyped: None,
564                    histogram: None,
565                    timestamp_ms: None,
566                },
567                Metric {
568                    label: vec![
569                        LabelPair {
570                            name: Some("method".to_owned()),
571                            value: Some("get".to_owned()),
572                        },
573                        LabelPair {
574                            name: Some("code".to_owned()),
575                            value: Some("400".to_owned()),
576                        },
577                    ],
578                    gauge: Some(Gauge { value: Some(3.0) }),
579                    counter: None,
580                    summary: None,
581                    untyped: None,
582                    histogram: None,
583                    timestamp_ms: None,
584                },
585            ],
586            unit: None,
587        };
588        let mut expected_msg = BytesMut::new();
589        expected.encode_length_delimited(&mut expected_msg).unwrap();
590
591        assert_eq!(actual_msg, expected_msg);
592
593        let actual = MetricFamily::decode_length_delimited(actual_msg).unwrap();
594        assert_eq!(actual, expected);
595    }
596}