diff --git a/benches/encoding/text.rs b/benches/encoding/text.rs index a270662..24a3b6c 100644 --- a/benches/encoding/text.rs +++ b/benches/encoding/text.rs @@ -10,7 +10,17 @@ use std::fmt::Write; use std::hint::black_box; pub fn text(c: &mut Criterion) { - c.bench_function("encode", |b| { + bench_text(c, "encode", "200"); + bench_text( + c, + "encode_realistic_string_labels", + "checkout-api-7d9f4c8b6f-2xkpt", + ); + bench_text(c, "encode_escaped_label_values", "2\\0\"0\n"); +} + +fn bench_text(c: &mut Criterion, name: &str, status_value: &'static str) { + c.bench_function(name, |b| { #[derive(Clone, Hash, PartialEq, Eq, EncodeLabelSet, Debug)] struct Labels { method: Method, @@ -26,23 +36,11 @@ pub fn text(c: &mut Criterion) { } #[derive(Clone, Hash, PartialEq, Eq, Debug)] - enum Status { - Two, - #[allow(dead_code)] - Four, - #[allow(dead_code)] - Five, - } + struct Status(&'static str); impl prometheus_client::encoding::EncodeLabelValue for Status { fn encode(&self, writer: &mut LabelValueEncoder) -> Result<(), std::fmt::Error> { - let status = match self { - Status::Two => "200", - Status::Four => "400", - Status::Five => "500", - }; - writer.write_str(status)?; - Ok(()) + writer.write_str(self.0) } } @@ -69,14 +67,14 @@ pub fn text(c: &mut Criterion) { counter_family .get_or_create(&Labels { method: Method::Get, - status: Status::Two, + status: Status(status_value), some_number: j.into(), }) .inc(); histogram_family .get_or_create(&Labels { method: Method::Get, - status: Status::Two, + status: Status(status_value), some_number: j.into(), }) .observe(j.into()); @@ -86,6 +84,7 @@ pub fn text(c: &mut Criterion) { let mut buffer = String::new(); b.iter(|| { + buffer.clear(); encoding::text::encode(&mut buffer, ®istry).unwrap(); black_box(&mut buffer); }) diff --git a/src/encoding.rs b/src/encoding.rs index 8b47d09..e3398f6 100644 --- a/src/encoding.rs +++ b/src/encoding.rs @@ -613,6 +613,10 @@ impl std::fmt::Write for LabelValueEncoder<'_> { } impl LabelValueEncoder<'_> { + pub(crate) fn write_str_unescaped(&mut self, s: &str) -> Result<(), std::fmt::Error> { + for_both_mut!(self, LabelValueEncoderInner, e, e.write_str_unescaped(s)) + } + /// Finish encoding the label value. pub fn finish(self) -> Result<(), std::fmt::Error> { for_both!(self, LabelValueEncoderInner, e, e.finish()) @@ -676,7 +680,7 @@ where impl EncodeLabelValue for f64 { fn encode(&self, encoder: &mut LabelValueEncoder) -> Result<(), std::fmt::Error> { - encoder.write_str(dtoa::Buffer::new().format(*self)) + encoder.write_str_unescaped(dtoa::Buffer::new().format(*self)) } } @@ -694,7 +698,7 @@ where impl EncodeLabelValue for bool { fn encode(&self, encoder: &mut LabelValueEncoder) -> Result<(), std::fmt::Error> { - encoder.write_str(if *self { "true" } else { "false" }) + encoder.write_str_unescaped(if *self { "true" } else { "false" }) } } @@ -702,7 +706,7 @@ macro_rules! impl_encode_label_value_for_integer { ($($t:ident),*) => {$( impl EncodeLabelValue for $t { fn encode(&self, encoder: &mut LabelValueEncoder) -> Result<(), std::fmt::Error> { - encoder.write_str(itoa::Buffer::new().format(*self)) + encoder.write_str_unescaped(itoa::Buffer::new().format(*self)) } } )*}; diff --git a/src/encoding/openmetrics_protobuf.rs b/src/encoding/openmetrics_protobuf.rs index d33e060..4280cb2 100644 --- a/src/encoding/openmetrics_protobuf.rs +++ b/src/encoding/openmetrics_protobuf.rs @@ -453,6 +453,11 @@ impl LabelValueEncoder<'_> { pub fn finish(self) -> Result<(), std::fmt::Error> { Ok(()) } + + pub(crate) fn write_str_unescaped(&mut self, s: &str) -> Result<(), std::fmt::Error> { + self.label_value.push_str(s); + Ok(()) + } } impl std::fmt::Write for LabelValueEncoder<'_> { diff --git a/src/encoding/prometheus_protobuf.rs b/src/encoding/prometheus_protobuf.rs index 621dbfc..3c11187 100644 --- a/src/encoding/prometheus_protobuf.rs +++ b/src/encoding/prometheus_protobuf.rs @@ -582,6 +582,11 @@ impl LabelValueEncoder<'_> { pub fn finish(self) -> Result<(), std::fmt::Error> { Ok(()) } + + pub(crate) fn write_str_unescaped(&mut self, s: &str) -> Result<(), std::fmt::Error> { + self.label_value.push_str(s); + Ok(()) + } } impl std::fmt::Write for LabelValueEncoder<'_> { diff --git a/src/encoding/text.rs b/src/encoding/text.rs index 4b8bd0b..fe02e3a 100644 --- a/src/encoding/text.rs +++ b/src/encoding/text.rs @@ -748,11 +748,39 @@ impl LabelValueEncoder<'_> { pub fn finish(self) -> Result<(), std::fmt::Error> { self.writer.write_str("\"") } + + pub(crate) fn write_str_unescaped(&mut self, s: &str) -> Result<(), std::fmt::Error> { + self.writer.write_str(s) + } + + #[cold] + fn write_str_escaped(&mut self, s: &str) -> Result<(), std::fmt::Error> { + let mut last = 0; + for (index, byte) in s.bytes().enumerate() { + let escaped = match byte { + b'\\' => "\\\\", + b'"' => "\\\"", + b'\n' => "\\n", + _ => continue, + }; + self.writer.write_str(&s[last..index])?; + self.writer.write_str(escaped)?; + last = index + 1; + } + self.writer.write_str(&s[last..]) + } } impl std::fmt::Write for LabelValueEncoder<'_> { fn write_str(&mut self, s: &str) -> std::fmt::Result { - self.writer.write_str(s) + let needs_escaping = s.bytes().fold(false, |found, byte| { + found | matches!(byte, b'\\' | b'"' | b'\n') + }); + if needs_escaping { + self.write_str_escaped(s) + } else { + self.write_str_unescaped(s) + } } } @@ -766,11 +794,102 @@ mod tests { use crate::metrics::info::Info; use crate::metrics::{counter::Counter, exemplar::CounterWithExemplar}; use pyo3::{prelude::*, types::PyModule}; + use quickcheck::QuickCheck; use std::borrow::Cow; use std::fmt::Error; use std::sync::atomic::{AtomicI32, AtomicU32}; use std::time::{SystemTime, UNIX_EPOCH}; + #[test] + fn label_values_escape_special_characters() { + let mut encoded = String::new(); + let mut label_set_encoder = LabelSetEncoder::new(&mut encoded); + let mut label_encoder = label_set_encoder.encode_label(); + let mut key_encoder = label_encoder.encode_label_key().unwrap(); + key_encoder.write_str("label").unwrap(); + let mut value_encoder = key_encoder.encode_label_value().unwrap(); + value_encoder + .write_str("plain \\ quoted \" line\ncarriage\r unicode λ") + .unwrap(); + value_encoder.finish().unwrap(); + + assert_eq!( + concat!( + r#"label="plain \\ quoted \" line\ncarriage"#, + "\r", + r#" unicode λ""# + ), + encoded + ); + } + + #[test] + fn label_value_escaping_matches_reference() { + fn reference(s: &str) -> String { + let mut escaped = String::new(); + for character in s.chars() { + match character { + '\\' => escaped.push_str("\\\\"), + '"' => escaped.push_str("\\\""), + '\n' => escaped.push_str("\\n"), + _ => escaped.push(character), + } + } + escaped + } + + fn encode(value: &str) -> String { + let mut encoded = String::new(); + let mut encoder = LabelValueEncoder { + writer: &mut encoded, + }; + encoder.write_str(value).unwrap(); + encoded + } + + fn prop(value: String) -> bool { + if encode(&value) != reference(&value) { + return false; + } + + // Guarantee that QuickCheck exercises a newline without another special + // byte that could independently select the escaping path. + let newline_only = format!("{}\n", value.replace(['\\', '"', '\n'], "")); + encode(&newline_only) == reference(&newline_only) + } + + QuickCheck::new() + .tests(1_000) + .quickcheck(prop as fn(String) -> bool); + } + + #[test] + fn escaped_label_values_produce_parseable_exposition() { + let mut registry = Registry::default(); + let family = Family::, Counter>::default(); + registry.register("requests", "Requests", family.clone()); + family + .get_or_create(&vec![( + "client_version".to_string(), + "a\"} evil{x=\"1\\line\nnext".to_string(), + )]) + .inc(); + + let mut encoded = String::new(); + encode(&mut encoded, ®istry).unwrap(); + + assert_eq!( + concat!( + "# HELP requests Requests.\n", + "# TYPE requests counter\n", + r#"requests_total{client_version="a\"} evil{x=\"1\\line\nnext"} 1"#, + "\n# EOF\n" + ), + encoded + ); + parse_with_python_client(encoded); + } + #[test] fn encode_counter() { let counter: Counter = Counter::default();