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
20pub struct ProtoEncoder<W> {
22 state: State,
23 pub writer: W,
24 buf: Vec<u8>,
25}
26
27impl<W: Write> ProtoEncoder<W> {
28 pub fn new(w: W) -> Self {
32 Self {
33 state: State::Init,
34 writer: w,
35 buf: Vec::new(),
36 }
37 }
38
39 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#[derive(Clone, Copy, Debug)]
73pub enum MetricType {
74 Counter,
76 Histogram,
78 Gauge,
80 Summary,
82 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 fn write_help(
93 &mut self,
94 name: impl MetricNameEncoder,
95 help: &str,
96 ) -> Result<(), std::io::Error> {
97 self.flush_buf()?;
98
99 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 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 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 encode_message(1, label_pair_len, self.buf, |buf| {
215 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 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 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 encode_message(4, metric_len, &mut enc.buf, |buf| {
265 labels.visit_values(&mut GroupVisitor { buf });
266
267 encode_message(3, count_len, buf, |buf| {
269 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 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 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 encode_message(4, metric_len, &mut enc.buf, |buf| {
320 labels.visit_values(&mut GroupVisitor { buf });
321
322 encode_message(2, gauge_len, buf, |buf| {
324 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 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 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 encode_message(4, metric_len, &mut enc.buf, |buf| {
375 labels.visit_values(&mut GroupVisitor { buf });
376
377 encode_message(2, gauge_len, buf, |buf| {
379 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}