Skip to main content

servers/prom_remote_write/
decode.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Decoding of Prometheus remote write protobuf payloads.
16
17use std::collections::BTreeMap;
18use std::collections::btree_map::Entry;
19
20use api::prom_store::remote::Sample;
21use bytes::Buf;
22use common_query::prelude::{greptime_timestamp, greptime_value};
23use pipeline::{ContextReq, GreptimePipelineParams, PipelineContext, PipelineDefinition};
24use prost::DecodeError;
25use prost::encoding::message::merge;
26use prost::encoding::{WireType, decode_key, decode_varint};
27use session::context::QueryContextRef;
28use snafu::OptionExt;
29use vrl::prelude::NotNan;
30use vrl::value::{KeyString, Value as VrlValue};
31
32use crate::error::InternalSnafu;
33use crate::http::event::PipelineIngestRequest;
34use crate::pipeline::run_pipeline;
35use crate::prom_remote_write::row_builder::{PromCtx, TablesBuilder};
36use crate::prom_remote_write::types::PromLabel;
37use crate::prom_remote_write::validation::PromValidationMode;
38#[allow(deprecated)]
39use crate::prom_store::{
40    DATABASE_LABEL_ALT_BYTES, DATABASE_LABEL_BYTES, METRIC_NAME_LABEL_BYTES,
41    PHYSICAL_TABLE_LABEL_ALT_BYTES, PHYSICAL_TABLE_LABEL_BYTES, SCHEMA_LABEL_BYTES,
42};
43use crate::query_handler::PipelineHandlerRef;
44use crate::repeated_field::{Clear, RepeatedField};
45
46#[derive(Default, Debug)]
47pub(crate) struct PromTimeSeries {
48    pub(crate) table_name: String,
49    pub(crate) schema: Option<String>,
50    pub(crate) physical_table: Option<String>,
51
52    pub(crate) labels: RepeatedField<PromLabel>,
53    pub(crate) samples: RepeatedField<Sample>,
54}
55
56impl Clear for PromTimeSeries {
57    fn clear(&mut self) {
58        self.table_name.clear();
59        self.schema.clear();
60        self.physical_table.clear();
61        // Labels borrow from the request buffer, which may be replaced after this reset.
62        for label in self.labels.iter_mut() {
63            label.clear();
64        }
65        self.labels.clear();
66        self.samples.clear();
67    }
68}
69
70impl PromTimeSeries {
71    pub fn merge_field(
72        &mut self,
73        tag: u32,
74        wire_type: WireType,
75        buf: &mut &[u8],
76        prom_validation_mode: PromValidationMode,
77    ) -> Result<(), DecodeError> {
78        const STRUCT_NAME: &str = "PromTimeSeries";
79        match tag {
80            1u32 => {
81                let label = self.labels.push_default();
82
83                let len = decode_varint(buf).map_err(|mut error| {
84                    error.push(STRUCT_NAME, "labels");
85                    error
86                })?;
87                let remaining = buf.remaining();
88                if len > remaining as u64 {
89                    return Err(DecodeError::new("buffer underflow"));
90                }
91
92                let limit = remaining - len as usize;
93                while buf.remaining() > limit {
94                    let (tag, wire_type) = decode_key(buf)?;
95                    label.merge_field(tag, wire_type, buf)?;
96                }
97                if buf.remaining() != limit {
98                    return Err(DecodeError::new("delimited length exceeded"));
99                }
100
101                #[allow(deprecated)]
102                let is_special_label = match label.name {
103                    METRIC_NAME_LABEL_BYTES => {
104                        self.table_name = prom_validation_mode.decode_string(label.value)?;
105                        true
106                    }
107                    SCHEMA_LABEL_BYTES => {
108                        self.schema = Some(prom_validation_mode.decode_string(label.value)?);
109                        true
110                    }
111                    DATABASE_LABEL_BYTES | DATABASE_LABEL_ALT_BYTES => {
112                        if self.schema.is_none() {
113                            self.schema = Some(prom_validation_mode.decode_string(label.value)?);
114                        }
115                        true
116                    }
117                    PHYSICAL_TABLE_LABEL_BYTES | PHYSICAL_TABLE_LABEL_ALT_BYTES => {
118                        self.physical_table =
119                            Some(prom_validation_mode.decode_string(label.value)?);
120                        true
121                    }
122                    _ => false,
123                };
124                if is_special_label {
125                    label.clear();
126                    self.labels.truncate(self.labels.len() - 1);
127                }
128
129                Ok(())
130            }
131            2u32 => {
132                let sample = self.samples.push_default();
133                merge(WireType::LengthDelimited, sample, buf, Default::default()).map_err(
134                    |mut error| {
135                        error.push(STRUCT_NAME, "samples");
136                        error
137                    },
138                )?;
139                Ok(())
140            }
141            3u32 => prost::encoding::skip_field(wire_type, tag, buf, Default::default()),
142            4u32 => Err(DecodeError::new(
143                "remote write v1 native histogram ingestion is unsupported; use remote write v2",
144            )),
145            _ => prost::encoding::skip_field(wire_type, tag, buf, Default::default()),
146        }
147    }
148
149    fn add_to_table_data<'a>(
150        &mut self,
151        table_builders: &mut TablesBuilder<'a>,
152        prom_validation_mode: PromValidationMode,
153    ) -> Result<(), DecodeError> {
154        let label_num = self.labels.len();
155        let row_num = self.samples.len();
156
157        let prom_ctx = PromCtx {
158            schema: self.schema.take(),
159            physical_table: self.physical_table.take(),
160        };
161
162        let table_data = table_builders.get_or_create_table_builder(
163            prom_ctx,
164            std::mem::take(&mut self.table_name),
165            label_num,
166            row_num,
167        );
168        table_data.add_labels_and_samples(
169            self.labels.as_slice(),
170            self.samples.as_slice(),
171            prom_validation_mode,
172        )?;
173
174        Ok(())
175    }
176}
177
178#[derive(Default, Debug)]
179pub struct PromWriteRequest<'a> {
180    pub(crate) table_data: TablesBuilder<'a>,
181    series: PromTimeSeries,
182}
183
184impl<'a> Clear for PromWriteRequest<'a> {
185    fn clear(&mut self) {
186        self.series.clear();
187        self.table_data.clear();
188    }
189}
190
191impl<'a> PromWriteRequest<'a> {
192    pub fn as_row_insert_requests(&mut self) -> ContextReq {
193        self.table_data.as_insert_requests()
194    }
195
196    pub fn decode(
197        &mut self,
198        buf: Vec<u8>,
199        prom_validation_mode: PromValidationMode,
200        processor: &mut PromSeriesProcessor,
201    ) -> Result<(), DecodeError> {
202        const STRUCT_NAME: &str = "PromWriteRequest";
203        self.clear();
204        self.table_data.set_raw_data(buf);
205        let mut offset = 0;
206        while offset < self.table_data.raw_data.len() {
207            let mut should_add_to_table_data = false;
208            let mut decoded_timeseries = false;
209            {
210                let raw_data = &self.table_data.raw_data;
211                let buf = &mut &raw_data[offset..];
212                let (tag, wire_type) = decode_key(buf)?;
213                if wire_type != WireType::LengthDelimited {
214                    return Err(DecodeError::new(format!(
215                        "invalid wire type: {:?}",
216                        wire_type
217                    )));
218                }
219                match tag {
220                    1u32 => {
221                        let len = decode_varint(buf).map_err(|mut e| {
222                            e.push(STRUCT_NAME, "timeseries");
223                            e
224                        })?;
225                        let remaining = buf.remaining();
226                        if len > remaining as u64 {
227                            return Err(DecodeError::new("buffer underflow"));
228                        }
229
230                        let limit = remaining - len as usize;
231                        while buf.remaining() > limit {
232                            let (tag, wire_type) = decode_key(buf)?;
233                            self.series
234                                .merge_field(tag, wire_type, buf, prom_validation_mode)?;
235                        }
236                        if buf.remaining() != limit {
237                            return Err(DecodeError::new("delimited length exceeded"));
238                        }
239
240                        if processor.use_pipeline {
241                            processor.consume_series_to_pipeline_map(
242                                &mut self.series,
243                                prom_validation_mode,
244                            )?;
245                        } else {
246                            should_add_to_table_data = true;
247                        }
248
249                        decoded_timeseries = true;
250                    }
251                    3u32 => {
252                        prost::encoding::skip_field(wire_type, tag, buf, Default::default())?;
253                    }
254                    _ => prost::encoding::skip_field(wire_type, tag, buf, Default::default())?,
255                }
256                offset = raw_data.len() - buf.remaining();
257            }
258
259            if should_add_to_table_data {
260                self.series
261                    .add_to_table_data(&mut self.table_data, prom_validation_mode)?;
262            }
263
264            if decoded_timeseries {
265                self.series.clear();
266            }
267        }
268
269        Ok(())
270    }
271}
272
273/// Hook injected into the PromWriteRequest decoding process.
274pub struct PromSeriesProcessor {
275    pub(crate) use_pipeline: bool,
276    pub(crate) table_values: BTreeMap<String, Vec<VrlValue>>,
277
278    pub(crate) pipeline_handler: Option<PipelineHandlerRef>,
279    pub(crate) query_ctx: Option<QueryContextRef>,
280    pub(crate) pipeline_def: Option<PipelineDefinition>,
281}
282
283impl PromSeriesProcessor {
284    pub fn default_processor() -> Self {
285        Self {
286            use_pipeline: false,
287            table_values: BTreeMap::new(),
288            pipeline_handler: None,
289            query_ctx: None,
290            pipeline_def: None,
291        }
292    }
293
294    pub fn set_pipeline(
295        &mut self,
296        handler: PipelineHandlerRef,
297        query_ctx: QueryContextRef,
298        pipeline_def: PipelineDefinition,
299    ) {
300        self.use_pipeline = true;
301        self.pipeline_handler = Some(handler);
302        self.query_ctx = Some(query_ctx);
303        self.pipeline_def = Some(pipeline_def);
304    }
305
306    pub(crate) fn consume_series_to_pipeline_map(
307        &mut self,
308        series: &mut PromTimeSeries,
309        prom_validation_mode: PromValidationMode,
310    ) -> Result<(), DecodeError> {
311        let mut vec_pipeline_map = Vec::new();
312        let mut pipeline_map = BTreeMap::new();
313        for l in series.labels.iter() {
314            let name = prom_validation_mode.decode_label_name(l.name)?;
315            let value = prom_validation_mode.decode_string(l.value)?;
316            pipeline_map.insert(KeyString::from(name), VrlValue::Bytes(value.into()));
317        }
318
319        let one_sample = series.samples.len() == 1;
320
321        for s in series.samples.iter() {
322            let Ok(value) = NotNan::new(s.value) else {
323                common_telemetry::warn!("Invalid float value: {}", s.value);
324                continue;
325            };
326
327            let timestamp = s.timestamp;
328            pipeline_map.insert(
329                KeyString::from(greptime_timestamp()),
330                VrlValue::Integer(timestamp),
331            );
332            pipeline_map.insert(KeyString::from(greptime_value()), VrlValue::Float(value));
333            if one_sample {
334                vec_pipeline_map.push(VrlValue::Object(pipeline_map));
335                break;
336            } else {
337                vec_pipeline_map.push(VrlValue::Object(pipeline_map.clone()));
338            }
339        }
340
341        let table_name = std::mem::take(&mut series.table_name);
342        match self.table_values.entry(table_name) {
343            Entry::Occupied(mut occupied_entry) => {
344                occupied_entry.get_mut().append(&mut vec_pipeline_map);
345            }
346            Entry::Vacant(vacant_entry) => {
347                vacant_entry.insert(vec_pipeline_map);
348            }
349        }
350
351        Ok(())
352    }
353
354    pub(crate) async fn exec_pipeline(&mut self) -> crate::error::Result<ContextReq> {
355        let handler = self.pipeline_handler.as_ref().context(InternalSnafu {
356            err_msg: "pipeline handler is not set",
357        })?;
358        let pipeline_def = self.pipeline_def.as_ref().context(InternalSnafu {
359            err_msg: "pipeline definition is not set",
360        })?;
361        let pipeline_param = GreptimePipelineParams::default();
362        let query_ctx = self.query_ctx.as_ref().context(InternalSnafu {
363            err_msg: "query context is not set",
364        })?;
365
366        let pipeline_ctx = PipelineContext::new(pipeline_def, &pipeline_param, query_ctx.channel());
367
368        let mut req = ContextReq::default();
369        let table_values = std::mem::take(&mut self.table_values);
370        for (table_name, pipeline_maps) in table_values.into_iter() {
371            let pipeline_req = PipelineIngestRequest {
372                table: table_name,
373                values: pipeline_maps,
374            };
375            let row_req =
376                run_pipeline(handler, &pipeline_ctx, pipeline_req, query_ctx, true).await?;
377            req.merge(row_req);
378        }
379
380        Ok(req)
381    }
382}
383
384#[cfg(test)]
385mod tests {
386    use std::collections::HashMap;
387
388    use api::prom_store::remote::{Histogram, Label, Sample, TimeSeries, WriteRequest};
389    use api::v1::{Row, RowInsertRequests, Rows};
390    use bytes::Bytes;
391    use prost::Message;
392
393    use super::*;
394    use crate::prom_store::to_grpc_row_insert_requests;
395    use crate::repeated_field::Clear;
396
397    fn sort_rows(rows: Rows) -> Rows {
398        let permutation =
399            permutation::sort_by_key(&rows.schema, |schema| schema.column_name.clone());
400        let schema = permutation.apply_slice(&rows.schema);
401        let mut inner_rows = vec![];
402        for row in rows.rows {
403            let values = permutation.apply_slice(&row.values);
404            inner_rows.push(Row { values });
405        }
406        Rows {
407            schema,
408            rows: inner_rows,
409        }
410    }
411
412    fn check_deserialized(
413        prom_write_request: &mut PromWriteRequest,
414        data: &[u8],
415        expected_samples: usize,
416        expected_rows: &RowInsertRequests,
417    ) {
418        let mut p = PromSeriesProcessor::default_processor();
419        prom_write_request.clear();
420        prom_write_request
421            .decode(data.to_owned(), PromValidationMode::Strict, &mut p)
422            .unwrap();
423
424        let req = prom_write_request.as_row_insert_requests();
425
426        let samples = req
427            .ref_all_req()
428            .filter_map(|r| r.rows.as_ref().map(|r| r.rows.len()))
429            .sum::<usize>();
430        let prom_rows = RowInsertRequests {
431            inserts: req.all_req().collect::<Vec<_>>(),
432        };
433
434        assert_eq!(expected_samples, samples);
435        assert_eq!(expected_rows.inserts.len(), prom_rows.inserts.len());
436
437        let expected_rows_map = expected_rows
438            .inserts
439            .iter()
440            .map(|insert| (insert.table_name.clone(), insert.rows.clone().unwrap()))
441            .collect::<HashMap<_, _>>();
442
443        for r in &prom_rows.inserts {
444            let expected_rows = expected_rows_map.get(&r.table_name).unwrap().clone();
445            assert_eq!(sort_rows(expected_rows), sort_rows(r.rows.clone().unwrap()));
446        }
447    }
448
449    #[test]
450    fn test_decode_write_request() {
451        let mut d = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"));
452        d.push("benches");
453        d.push("write_request.pb.data");
454        let data = std::fs::read(d).unwrap();
455
456        let (expected_rows, expected_samples) =
457            to_grpc_row_insert_requests(&WriteRequest::decode(&data[..]).unwrap()).unwrap();
458
459        let mut prom_write_request = PromWriteRequest::default();
460        for _ in 0..3 {
461            check_deserialized(
462                &mut prom_write_request,
463                &data,
464                expected_samples,
465                &expected_rows,
466            );
467        }
468    }
469
470    #[test]
471    fn test_decode_rejects_remote_write_v1_native_histograms() {
472        for samples in [
473            Vec::new(),
474            vec![Sample {
475                value: 1.0,
476                timestamp: 1000,
477            }],
478        ] {
479            let request = WriteRequest {
480                timeseries: vec![TimeSeries {
481                    samples,
482                    histograms: vec![Histogram::default()],
483                    ..Default::default()
484                }],
485                ..Default::default()
486            };
487            let mut processor = PromSeriesProcessor::default_processor();
488            let mut write_request = PromWriteRequest::default();
489            let error = write_request
490                .decode(
491                    request.encode_to_vec(),
492                    PromValidationMode::Strict,
493                    &mut processor,
494                )
495                .unwrap_err();
496
497            assert!(error.to_string().contains(
498                "remote write v1 native histogram ingestion is unsupported; use remote write v2"
499            ));
500            assert_eq!(
501                write_request.as_row_insert_requests().ref_all_req().count(),
502                0
503            );
504        }
505    }
506
507    #[test]
508    fn test_decode_clears_state_after_error() {
509        let label = |name: &str, value: &str| Label {
510            name: name.to_string(),
511            value: value.to_string(),
512        };
513        let sample = Sample {
514            value: 1.0,
515            timestamp: 1000,
516        };
517        let failed_request = WriteRequest {
518            timeseries: vec![
519                TimeSeries {
520                    labels: vec![
521                        label("__name__", "stale_metric"),
522                        label("stale_label", "stale_value"),
523                    ],
524                    samples: vec![sample],
525                    ..Default::default()
526                },
527                TimeSeries {
528                    labels: vec![
529                        label("__name__", "rejected_metric"),
530                        label("rejected_label", "rejected_value"),
531                        label("__schema__", "rejected_schema"),
532                        label("x_greptime_physical_table", "rejected_physical_table"),
533                    ],
534                    samples: vec![sample],
535                    histograms: vec![Histogram::default()],
536                    ..Default::default()
537                },
538            ],
539            ..Default::default()
540        };
541        let successful_request = WriteRequest {
542            timeseries: vec![TimeSeries {
543                labels: vec![
544                    label("__name__", "fresh_metric"),
545                    label("fresh_label", "fresh_value"),
546                ],
547                samples: vec![sample],
548                ..Default::default()
549            }],
550            ..Default::default()
551        };
552        let mut processor = PromSeriesProcessor::default_processor();
553        let mut write_request = PromWriteRequest::default();
554
555        write_request
556            .decode(
557                failed_request.encode_to_vec(),
558                PromValidationMode::Strict,
559                &mut processor,
560            )
561            .unwrap_err();
562        write_request
563            .decode(
564                successful_request.encode_to_vec(),
565                PromValidationMode::Strict,
566                &mut processor,
567            )
568            .unwrap();
569
570        assert_eq!(write_request.table_data.tables.len(), 1);
571        let (prom_ctx, tables) = write_request.table_data.tables.iter().next().unwrap();
572        assert_eq!(prom_ctx.schema, None);
573        assert_eq!(prom_ctx.physical_table, None);
574        assert_eq!(tables.len(), 1);
575        assert!(tables.contains_key("fresh_metric"));
576
577        let requests = write_request
578            .as_row_insert_requests()
579            .all_req()
580            .collect::<Vec<_>>();
581        assert_eq!(requests.len(), 1);
582        assert_eq!(requests[0].table_name, "fresh_metric");
583        let rows = requests[0].rows.as_ref().unwrap();
584        assert_eq!(rows.rows.len(), 1);
585        assert_eq!(
586            rows.schema
587                .iter()
588                .map(|column| column.column_name.as_str())
589                .collect::<Vec<_>>(),
590            vec![greptime_timestamp(), greptime_value(), "fresh_label"]
591        );
592    }
593
594    #[test]
595    fn test_decode_string_strict_mode_valid_utf8() {
596        let valid_utf8 = Bytes::from("hello world");
597        let result = PromValidationMode::Strict.decode_string(&valid_utf8);
598        assert!(result.is_ok());
599        assert_eq!(result.unwrap(), "hello world");
600    }
601
602    #[test]
603    fn test_decode_string_all_modes_ascii() {
604        let ascii = Bytes::from("simple_ascii_123");
605        let strict_result = PromValidationMode::Strict.decode_string(&ascii).unwrap();
606        let lossy_result = PromValidationMode::Lossy.decode_string(&ascii).unwrap();
607        let unchecked_result = PromValidationMode::Unchecked.decode_string(&ascii).unwrap();
608        assert_eq!(strict_result, "simple_ascii_123");
609        assert_eq!(lossy_result, "simple_ascii_123");
610        assert_eq!(unchecked_result, "simple_ascii_123");
611    }
612}