1use 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 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
273pub 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}