diff --git a/src/servers/Cargo.toml b/src/servers/Cargo.toml index 90dd189022..d64a1bd5a5 100644 --- a/src/servers/Cargo.toml +++ b/src/servers/Cargo.toml @@ -180,6 +180,11 @@ common-version.workspace = true name = "bench_prom" harness = false +[[bench]] +name = "prom_decode" +harness = false +required-features = ["testing"] + [[bench]] name = "to_http_output" harness = false diff --git a/src/servers/benches/prom_decode.rs b/src/servers/benches/prom_decode.rs index e9db0ad771..06b0d02e94 100644 --- a/src/servers/benches/prom_decode.rs +++ b/src/servers/benches/prom_decode.rs @@ -12,21 +12,137 @@ // See the License for the specific language governing permissions and // limitations under the License. +use std::collections::HashMap; use std::time::Duration; -use api::prom_store::remote::WriteRequest; -use criterion::{BenchmarkId, Criterion, black_box, criterion_group, criterion_main}; +use api::greptime_proto::io::prometheus::write::v2 as write_v2; +use api::prom_store::remote::{self as write_v1, WriteRequest}; +use criterion::{ + BatchSize, BenchmarkId, Criterion, Throughput, black_box, criterion_group, criterion_main, +}; use prost::Message; use servers::prom_remote_write::decode::{PromSeriesProcessor, PromWriteRequest}; +use servers::prom_remote_write::v2::test_util as remote_write_v2; use servers::prom_remote_write::validation::{PromValidationMode, validate_label_name}; use servers::prom_store::to_grpc_row_insert_requests; -fn bench_decode_prom_request(c: &mut Criterion) { +fn load_fixture_v1_bytes() -> Vec { let mut d = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")); d.push("benches"); d.push("write_request.pb.data"); + std::fs::read(d).expect("read write_request.pb.data fixture") +} - let data = std::fs::read(d).unwrap(); +/// Convert a PRW v1 `WriteRequest` into an equivalent PRW v2 `Request`. +/// +/// Labels are interned into the v2 symbol table so payloads with repeated label +/// names/values exercise the main on-wire advantage of v2. +fn write_request_v1_to_v2(v1: &WriteRequest) -> write_v2::Request { + let mut symbols = vec![String::new()]; + let mut symbol_index: HashMap = HashMap::new(); + symbol_index.insert(String::new(), 0); + + let mut intern = |value: &str| -> u32 { + if let Some(&idx) = symbol_index.get(value) { + return idx; + } + let idx = symbols.len() as u32; + symbols.push(value.to_string()); + symbol_index.insert(value.to_string(), idx); + idx + }; + + let timeseries = v1 + .timeseries + .iter() + .map(|series| { + let mut labels: Vec<&write_v1::Label> = series.labels.iter().collect(); + // PRW v2 requires lexicographically sorted label pairs. + labels.sort_by(|a, b| a.name.cmp(&b.name).then_with(|| a.value.cmp(&b.value))); + + let mut labels_refs = Vec::with_capacity(labels.len() * 2); + for label in labels { + labels_refs.push(intern(&label.name)); + labels_refs.push(intern(&label.value)); + } + + write_v2::TimeSeries { + labels_refs, + samples: series + .samples + .iter() + .map(|sample| write_v2::Sample { + value: sample.value, + timestamp: sample.timestamp, + start_timestamp: 0, + }) + .collect(), + histograms: Vec::new(), + exemplars: Vec::new(), + metadata: None, + } + }) + .collect(); + + write_v2::Request { + symbols, + timeseries, + } +} + +/// Synthetic sample workload with controllable label reuse. +/// +/// `shared_labels` keeps common label names/values across series (v2-friendly). +/// Each series also gets a unique `series_id` label so cardinality stays high. +fn generate_sample_workload( + series_count: usize, + samples_per_series: usize, + shared_label_count: usize, +) -> (WriteRequest, write_v2::Request) { + let shared_labels: Vec = (0..shared_label_count) + .map(|i| write_v1::Label { + name: format!("label_{i}"), + value: format!("value_{}", i % 16), + }) + .collect(); + + let timeseries = (0..series_count) + .map(|series_idx| { + let mut labels = Vec::with_capacity(shared_label_count + 2); + labels.push(write_v1::Label { + name: "__name__".to_string(), + value: format!("metric_{}", series_idx % 32), + }); + labels.extend(shared_labels.iter().cloned()); + labels.push(write_v1::Label { + name: "series_id".to_string(), + value: series_idx.to_string(), + }); + + write_v1::TimeSeries { + labels, + samples: (0..samples_per_series) + .map(|sample_idx| write_v1::Sample { + value: (series_idx * 1000 + sample_idx) as f64, + timestamp: 1_700_000_000_000 + sample_idx as i64 * 1_000, + }) + .collect(), + exemplars: Vec::new(), + histograms: Vec::new(), + } + }) + .collect(); + + let v1 = WriteRequest { + timeseries, + metadata: Vec::new(), + }; + let v2 = write_request_v1_to_v2(&v1); + (v1, v2) +} + +fn bench_decode_prom_request(c: &mut Criterion) { + let data = load_fixture_v1_bytes(); let mut group = c.benchmark_group("decode_prom_request"); group.measurement_time(Duration::from_secs(3)); @@ -64,6 +180,128 @@ fn bench_decode_prom_request(c: &mut Criterion) { group.finish(); } +/// Compare Prometheus remote-write v1 vs v2 decoding efficiency. +/// +/// Two layers are measured on equivalent logical payloads: +/// 1. protobuf decode only (`WriteRequest` / v2 `Request`) +/// 2. full Greptime path to row-insert requests +/// +/// Workloads: +/// - `fixture`: existing `write_request.pb.data` converted v1 -> v2 +/// - `synthetic_*`: generated series with shared labels (symbol-table friendly) +#[allow(clippy::print_stderr)] +fn bench_prom_v1_vs_v2_decode(c: &mut Criterion) { + let fixture_v1_bytes = load_fixture_v1_bytes(); + let fixture_v1 = WriteRequest::decode(fixture_v1_bytes.as_slice()).unwrap(); + let fixture_v2 = write_request_v1_to_v2(&fixture_v1); + let fixture_v2_bytes = fixture_v2.encode_to_vec(); + + let synthetic = [ + ( + "synthetic_1k_series_1_sample", + generate_sample_workload(1_000, 1, 8), + ), + ( + "synthetic_1k_series_10_samples", + generate_sample_workload(1_000, 10, 8), + ), + ( + "synthetic_5k_series_1_sample", + generate_sample_workload(5_000, 1, 8), + ), + ]; + + let mut workloads: Vec<(&str, Vec, Vec)> = + vec![("fixture", fixture_v1_bytes, fixture_v2_bytes)]; + for (name, (v1, v2)) in &synthetic { + workloads.push((*name, v1.encode_to_vec(), v2.encode_to_vec())); + } + + eprintln!("\nprom remote-write v1 vs v2 payload sizes:"); + for (name, v1_bytes, v2_bytes) in &workloads { + let ratio = v2_bytes.len() as f64 / v1_bytes.len() as f64; + eprintln!( + " {name}: v1={} B, v2={} B, v2/v1={ratio:.3}", + v1_bytes.len(), + v2_bytes.len() + ); + } + + // --- protobuf decode only --- + { + let mut group = c.benchmark_group("prom_rw_protobuf_decode"); + group.measurement_time(Duration::from_secs(3)); + + for (name, v1_bytes, v2_bytes) in &workloads { + group.throughput(Throughput::Bytes(v1_bytes.len() as u64)); + group.bench_with_input(BenchmarkId::new("v1", name), v1_bytes, |b, bytes| { + b.iter(|| { + black_box(WriteRequest::decode(black_box(bytes.as_slice())).unwrap()); + }); + }); + + group.throughput(Throughput::Bytes(v2_bytes.len() as u64)); + group.bench_with_input(BenchmarkId::new("v2", name), v2_bytes, |b, bytes| { + b.iter(|| { + black_box(write_v2::Request::decode(black_box(bytes.as_slice())).unwrap()); + }); + }); + } + + group.finish(); + } + + // --- full path: protobuf -> Greptime row inserts --- + { + let mut group = c.benchmark_group("prom_rw_decode_to_rows"); + group.sample_size(50); + group.measurement_time(Duration::from_secs(5)); + + for (name, v1_bytes, v2_bytes) in &workloads { + group.throughput(Throughput::Bytes(v1_bytes.len() as u64)); + group.bench_with_input(BenchmarkId::new("v1_custom", name), v1_bytes, |b, bytes| { + let mut prom_request = PromWriteRequest::default(); + let mut processor = PromSeriesProcessor::default_processor(); + b.iter_batched( + || bytes.clone(), + |bytes| { + prom_request + .decode(black_box(bytes), PromValidationMode::Strict, &mut processor) + .unwrap(); + let rows = prom_request.as_row_insert_requests(); + black_box(&rows); + }, + BatchSize::LargeInput, + ); + }); + + // Baseline: stock prost WriteRequest + existing converter. + group.throughput(Throughput::Bytes(v1_bytes.len() as u64)); + group.bench_with_input(BenchmarkId::new("v1_prost", name), v1_bytes, |b, bytes| { + b.iter(|| { + let request = WriteRequest::decode(black_box(bytes.as_slice())).unwrap(); + black_box(to_grpc_row_insert_requests(&request).unwrap()); + }); + }); + + group.throughput(Throughput::Bytes(v2_bytes.len() as u64)); + group.bench_with_input(BenchmarkId::new("v2_manual", name), v2_bytes, |b, bytes| { + b.iter(|| { + black_box( + remote_write_v2::decode_uncompressed_write_requests( + black_box(bytes.as_slice()), + true, + ) + .unwrap(), + ); + }); + }); + } + + group.finish(); + } +} + /// Benchmark comparing UTF-8 string validation (`decode_string`) vs /// direct byte-level Prometheus label name validation (`decode_label_name`). fn bench_label_name_validation(c: &mut Criterion) { @@ -155,6 +393,7 @@ fn bench_utf8_validation(c: &mut Criterion) { criterion_group!( benches, bench_decode_prom_request, + bench_prom_v1_vs_v2_decode, bench_label_name_validation, bench_utf8_validation ); diff --git a/src/servers/src/http/prom_store.rs b/src/servers/src/http/prom_store.rs index ac4f39628d..ed960264b0 100644 --- a/src/servers/src/http/prom_store.rs +++ b/src/servers/src/http/prom_store.rs @@ -48,7 +48,7 @@ use crate::http::header::{ use crate::pending_rows_batcher::PendingRowsBatcher; use crate::prom_remote_write::decode::PromSeriesProcessor; use crate::prom_remote_write::decode_remote_write_request; -use crate::prom_remote_write::v2::{decode_remote_write_v2_request, into_write_requests}; +use crate::prom_remote_write::v2::decode_remote_write_v2; use crate::prom_remote_write::validation::PromValidationMode; use crate::prom_store::snappy_decompress; use crate::query_handler::{PipelineHandlerRef, PromStoreProtocolHandlerRef, PromStoreResponse}; @@ -235,23 +235,11 @@ async fn remote_write_v2( let (db, mut query_ctx, _timer) = prepare_remote_write_context(¶ms, query_ctx, REMOTE_WRITE_V2_VERSION); - let request = match decode_remote_write_v2_request(is_zstd, body) { - Ok(request) => request, - Err(error) => return Ok(remote_write_v2_error_response(error, 0, 0, 0)), - }; - if !experimental_enable_prometheus_native_histogram && request_has_native_histograms(&request) { - return Ok(remote_write_v2_error_response( - error::InvalidPromRemoteRequestSnafu { - msg: "prometheus remote write v2 native histogram ingestion is experimental; set prom_store.experimental_enable_prometheus_native_histogram = true to enable it" - .to_string(), - } - .build(), - 0, - 0, - 0, - )); - } - let req = match into_write_requests(request) { + let req = match decode_remote_write_v2( + is_zstd, + body, + experimental_enable_prometheus_native_histogram, + ) { Ok(req) => req, Err(error) => return Ok(remote_write_v2_error_response(error, 0, 0, 0)), }; @@ -302,15 +290,6 @@ async fn remote_write_v2( Ok((StatusCode::NO_CONTENT, headers).into_response()) } -fn request_has_native_histograms( - request: &api::greptime_proto::io::prometheus::write::v2::Request, -) -> bool { - request - .timeseries - .iter() - .any(|series| !series.histograms.is_empty()) -} - fn vm_proto_version_response(params: &RemoteWriteQuery) -> Option { params .get_vm_proto_version diff --git a/src/servers/src/prom_remote_write/README.md b/src/servers/src/prom_remote_write/README.md index d21eead190..9e86a8753d 100644 --- a/src/servers/src/prom_remote_write/README.md +++ b/src/servers/src/prom_remote_write/README.md @@ -11,28 +11,33 @@ Remote write v2 enters through `remote_write_v2` in ```mermaid flowchart TD A["HTTP /v1/prometheus/write"] --> B["remote_write_v2"] - B --> C["decode_remote_write_v2_request"] - C --> D["into_write_requests"] + B --> C["decode_remote_write_v2"] + C --> D["borrow symbols and time-series bytes"] + D --> E["decode leaf messages and build rows"] - D --> E["samples ContextReq"] - D --> F["histograms ContextReq"] + E --> F["samples ContextReq"] + E --> G["histograms ContextReq"] - E --> G["write_prometheus_rows_with_progress"] - G --> H["metric engine / pending batcher when enabled"] + F --> H["write_prometheus_rows_with_progress"] + H --> I["metric engine / pending batcher when enabled"] - F --> I["histogram ContextOpt"] - I --> J["write_prometheus_rows_with_progress"] - J --> K["same metric-engine flag as samples, no batcher"] - K --> L["table: "] + G --> J["histogram ContextOpt"] + J --> K["write_prometheus_rows_with_progress"] + K --> L["same metric-engine flag as samples, no batcher"] + L --> M["table: "] - L --> M["field: configured native-histogram Struct"] - M --> N["struct children: counts, spans, buckets, sum, schema"] - H --> P["written headers and counters"] - K --> P + M --> N["field: configured native-histogram Struct"] + N --> O["struct children: counts, spans, buckets, sum, schema"] + I --> P["written headers and counters"] + L --> P ``` The conversion step splits one v2 request into two `ContextReq`s: +The `decode` codec metric covers decompression and the borrowed request +envelope. The `convert` metric covers per-series scans, leaf prost decoding, +validation, and row construction. + - samples keep the existing sample table name and can use the metric-engine physical table path and pending rows batcher; - native histograms keep the existing metric table name. diff --git a/src/servers/src/prom_remote_write/v2.rs b/src/servers/src/prom_remote_write/v2.rs index 5a25feb566..31da7ac1f3 100644 --- a/src/servers/src/prom_remote_write/v2.rs +++ b/src/servers/src/prom_remote_write/v2.rs @@ -17,19 +17,24 @@ use std::collections::hash_map::Entry; use ahash::{HashMap, HashMapExt, HashSet, HashSetExt}; use api::greptime_proto::io::prometheus::write::v2::histogram::{Count, ZeroCount}; use api::greptime_proto::io::prometheus::write::v2::{ - BucketSpan, Histogram, Request, Sample, TimeSeries, metadata, + BucketSpan, Exemplar, Histogram, Metadata, Sample, metadata, }; #[cfg(test)] -use api::greptime_proto::io::prometheus::write::v2::{Exemplar, Metadata}; +use api::greptime_proto::io::prometheus::write::v2::{Request, TimeSeries}; use api::helper::ColumnDataTypeWrapper; use api::v1::value::ValueData; -use api::v1::{ColumnSchema, ListValue, RowInsertRequest, Rows, SemanticType, Value}; -use bytes::Bytes; +use api::v1::{ + ColumnDataType, ColumnSchema, ListValue, RowInsertRequest, Rows, SemanticType, Value, +}; +use bytes::{Buf, Bytes}; use common_grpc::precision::Precision; use common_query::native_histogram::*; use common_query::prelude::{greptime_native_histogram, greptime_timestamp, greptime_value}; use pipeline::{ContextOpt, ContextReq}; -use prost::Message; +use prost::encoding::{ + DecodeContext, WireType, decode_key, decode_varint, message, skip_field, uint32, +}; +use prost::{DecodeError, Message}; use snafu::{OptionExt, ResultExt, ensure}; use table::requests::{ METADATA_QUALITY_DECLARED, SEMANTIC_METRIC_METADATA_QUALITY, SEMANTIC_METRIC_TYPE, @@ -52,23 +57,60 @@ use crate::semantic::{ openmetrics_unit_to_ucum, }; -type PromTags = Vec<(String, String)>; -type ResolvedSeriesLabels = (PromCtx, String, PromTags); +type PromTags<'a> = Vec<(&'a str, String)>; +type ResolvedSeriesLabels<'a> = (PromCtx, String, PromTags<'a>); const MAX_REMOTE_WRITE_V2_SCHEMA: i32 = 8; const MAX_REDUCIBLE_REMOTE_WRITE_V2_SCHEMA: i32 = 52; +const TIME_SERIES_LABELS_REFS_TAG: u32 = 1; +const TIME_SERIES_SAMPLES_TAG: u32 = 2; +const TIME_SERIES_HISTOGRAMS_TAG: u32 = 3; +const TIME_SERIES_EXEMPLARS_TAG: u32 = 4; +const TIME_SERIES_METADATA_TAG: u32 = 5; -pub(crate) fn decode_remote_write_v2_request(is_zstd: bool, body: Bytes) -> Result { - let _timer = crate::metrics::METRIC_HTTP_PROM_STORE_DECODE_ELAPSED.start_timer(); +struct BorrowedRequest<'a> { + symbols: Vec<&'a str>, + timeseries: Vec<&'a [u8]>, +} - // Match the v1 decoder's VictoriaMetrics fallback: some clients may send a - // mismatched content-encoding header, so try the other compression on failure. - let buf = if let Ok(buf) = try_decompress(is_zstd, &body[..]) { - buf - } else { - try_decompress(!is_zstd, &body[..])? - }; +impl<'a> BorrowedRequest<'a> { + fn decode(mut buf: &'a [u8]) -> std::result::Result { + let mut symbols = Vec::new(); + let mut timeseries = Vec::new(); - Request::decode(&buf[..]).context(error::DecodePromRemoteRequestSnafu) + while buf.has_remaining() { + let (tag, wire_type) = decode_key(&mut buf)?; + match tag { + 4 => { + let value = + take_length_delimited(wire_type, &mut buf).map_err(|mut error| { + error.push("Request", "symbols"); + error + })?; + let symbol = std::str::from_utf8(value).map_err(|_| { + let mut error = + DecodeError::new("invalid string value: data is not UTF-8 encoded"); + error.push("Request", "symbols"); + error + })?; + symbols.push(symbol); + } + 5 => { + let series = + take_length_delimited(wire_type, &mut buf).map_err(|mut error| { + error.push("Request", "timeseries"); + error + })?; + timeseries.push(series); + } + _ => skip_field(wire_type, tag, &mut buf, DecodeContext::default())?, + } + } + + Ok(Self { + symbols, + timeseries, + }) + } } pub(crate) struct RemoteWriteV2WriteRequests { @@ -81,19 +123,33 @@ pub(crate) struct RemoteWriteV2WriteRequests { pub semantic_index: SemanticIndexes, } -/// Converts a PRW v2 request into normal sample writes and native histogram writes. -/// -/// A metric name may appear in only one payload kind per request because the -/// metric-engine logical table is either a float metric or a native histogram. -pub(crate) fn into_write_requests(request: Request) -> Result { - let _timer = crate::metrics::METRIC_HTTP_PROM_STORE_CONVERT_ELAPSED.start_timer(); - let Request { - symbols, - timeseries, - } = request; +pub(crate) fn decode_remote_write_v2( + is_zstd: bool, + body: Bytes, + native_histograms_enabled: bool, +) -> Result { + let decode_timer = crate::metrics::METRIC_HTTP_PROM_STORE_DECODE_ELAPSED.start_timer(); + // Match the v1 decoder's VictoriaMetrics fallback: some clients may send a + // mismatched content-encoding header, so try the other compression on failure. + let buf = if let Ok(buf) = try_decompress(is_zstd, &body[..]) { + buf + } else { + try_decompress(!is_zstd, &body[..])? + }; + let request = BorrowedRequest::decode(&buf).context(error::DecodePromRemoteRequestSnafu)?; + drop(decode_timer); + + let _convert_timer = crate::metrics::METRIC_HTTP_PROM_STORE_CONVERT_ELAPSED.start_timer(); + convert_remote_write_v2(request, native_histograms_enabled) +} + +fn convert_remote_write_v2( + request: BorrowedRequest<'_>, + native_histograms_enabled: bool, +) -> Result { ensure!( - symbols.first().map(|s| s.as_str()) == Some(""), + request.symbols.first().copied() == Some(""), error::InvalidPromRemoteRequestSnafu { msg: "remote write v2 symbols must start with an empty string".to_string(), } @@ -104,28 +160,40 @@ pub(crate) fn into_write_requests(request: Request) -> Result 0 { - histogram_count > 0 + let has_other_value_type = if counts.samples > 0 { + counts.histograms > 0 || histogram_tables .get(&prom_ctx) .is_some_and(|tables| tables.contains_key(&table_name)) @@ -143,42 +211,40 @@ pub(crate) fn into_write_requests(request: Request) -> Result 0 && histogram_count == 0 { - // Fast path for regular sample-only series. Move the resolved labels into - // the sample writer instead of cloning them for a histogram path we won't use. - let table_data = get_or_create_table_data( - &mut sample_tables, - prom_ctx, - table_name, - tags.len() + 2, - sample_count, - ); + let column_count = + tags.len() + .checked_add(2) + .with_context(|| error::InvalidPromRemoteRequestSnafu { + msg: "remote write v2 series has too many labels".to_string(), + })?; + let (writer, row_count) = if counts.samples > 0 { + ( + SeriesWriter::Samples(get_or_create_table_data( + &mut sample_tables, + prom_ctx, + table_name, + column_count, + counts.samples, + )), + counts.samples, + ) + } else { + ( + SeriesWriter::Histograms(get_or_create_table_data( + &mut histogram_tables, + prom_ctx, + table_name, + column_count, + counts.histograms, + )), + counts.histograms, + ) + }; - write_samples(table_data, series.samples, tags)?; - sample_count_total += sample_count as u64; - // The owned labels were moved above, so skip the mixed-series path below. - continue; - } - - if histogram_count > 0 { - let table_data = get_or_create_table_data( - &mut histogram_tables, - prom_ctx, - table_name, - tags.len() + 2, - histogram_count, - ); - - let mut histograms = series.histograms; - let Some(last_histogram) = histograms.pop() else { - continue; - }; - for histogram in &histograms { - write_native_histogram(table_data, histogram, tags.iter().cloned())?; - } - write_native_histogram(table_data, &last_histogram, tags.into_iter())?; - histogram_count_total += histogram_count as u64; - } + decode_series_leaves(series, Some(writer), tags, row_count, &mut scratch)?; + sample_count_total = checked_total(sample_count_total, counts.samples, "sample")?; + histogram_count_total = + checked_total(histogram_count_total, counts.histograms, "histogram")?; } Ok(RemoteWriteV2WriteRequests { @@ -201,14 +267,11 @@ pub(crate) fn into_write_requests(request: Request) -> Result Result<()> { - let Some(series_metadata) = &series.metadata else { - return Ok(()); - }; // Symbol 0 is the mandatory empty string: no help / no unit. if series_metadata.help_ref != 0 { symbol_ref(symbols, series_metadata.help_ref, "metadata help")?; @@ -258,6 +321,256 @@ fn metric_type_value(wire_type: i32) -> Option<&'static str> { } } +#[derive(Default)] +struct SeriesCounts { + samples: usize, + histograms: usize, +} + +fn scan_series( + mut buf: &[u8], + labels_refs: &mut Vec, + metadata: &mut Metadata, +) -> std::result::Result { + labels_refs.clear(); + metadata.clear(); + let mut counts = SeriesCounts::default(); + + while buf.has_remaining() { + let (tag, wire_type) = decode_key(&mut buf)?; + match tag { + TIME_SERIES_LABELS_REFS_TAG => { + uint32::merge_repeated(wire_type, labels_refs, &mut buf, DecodeContext::default()) + .map_err(|mut error| { + error.push("TimeSeries", "labels_refs"); + error + })? + } + TIME_SERIES_SAMPLES_TAG => { + take_length_delimited(wire_type, &mut buf).map_err(|mut error| { + error.push("TimeSeries", "samples"); + error + })?; + counts.samples = counts.samples.checked_add(1).ok_or_else(|| { + DecodeError::new("remote write v2 sample count overflows usize") + })?; + } + TIME_SERIES_HISTOGRAMS_TAG => { + take_length_delimited(wire_type, &mut buf).map_err(|mut error| { + error.push("TimeSeries", "histograms"); + error + })?; + counts.histograms = counts.histograms.checked_add(1).ok_or_else(|| { + DecodeError::new("remote write v2 histogram count overflows usize") + })?; + } + TIME_SERIES_EXEMPLARS_TAG => { + take_length_delimited(wire_type, &mut buf)?; + } + TIME_SERIES_METADATA_TAG => { + message::merge(wire_type, metadata, &mut buf, DecodeContext::default()).map_err( + |mut error| { + error.push("TimeSeries", "metadata"); + error + }, + )? + } + _ => skip_field(wire_type, tag, &mut buf, DecodeContext::default())?, + } + } + + Ok(counts) +} + +#[derive(Default)] +struct LeafScratch { + sample: Sample, + histogram: Histogram, + exemplar: Exemplar, +} + +enum SeriesWriter<'a> { + Samples(&'a mut TableData), + Histograms(&'a mut TableData), +} + +fn decode_series_leaves( + mut buf: &[u8], + mut writer: Option>, + mut tags: PromTags<'_>, + mut rows_remaining: usize, + scratch: &mut LeafScratch, +) -> Result<()> { + let mut sample_row_template = None; + + while buf.has_remaining() { + let (tag, wire_type) = decode_key(&mut buf).context(error::DecodePromRemoteRequestSnafu)?; + match tag { + TIME_SERIES_LABELS_REFS_TAG => { + skip_field(wire_type, tag, &mut buf, DecodeContext::default()) + .context(error::DecodePromRemoteRequestSnafu)? + } + TIME_SERIES_SAMPLES_TAG => { + scratch.sample.clear(); + message::merge( + wire_type, + &mut scratch.sample, + &mut buf, + DecodeContext::default(), + ) + .map_err(|mut error| { + error.push("TimeSeries", "samples"); + error + }) + .context(error::DecodePromRemoteRequestSnafu)?; + rows_remaining = rows_remaining.checked_sub(1).with_context(|| { + error::InvalidPromRemoteRequestSnafu { + msg: "remote write v2 sample count changed between scans".to_string(), + } + })?; + if let Some(SeriesWriter::Samples(table_data)) = &mut writer { + if sample_row_template.is_none() { + let timestamp_index = table_data.ensure_column_by_name( + greptime_timestamp(), + ColumnDataType::TimestampMillisecond, + SemanticType::Timestamp, + )?; + let value_index = table_data.ensure_column_by_name( + greptime_value(), + ColumnDataType::Float64, + SemanticType::Field, + )?; + let mut row = table_data.alloc_one_row(); + row_writer::write_tags( + table_data, + std::mem::take(&mut tags).into_iter(), + &mut row, + )?; + sample_row_template = Some((row, timestamp_index, value_index)); + } + + if let Some((row, timestamp_index, value_index)) = sample_row_template.as_mut() + { + row[*timestamp_index].value_data = Some( + ValueData::TimestampMillisecondValue(scratch.sample.timestamp), + ); + row[*value_index].value_data = + Some(ValueData::F64Value(scratch.sample.value)); + let row = if rows_remaining == 0 { + std::mem::take(row) + } else { + row.clone() + }; + table_data.add_row(row); + } + } + } + TIME_SERIES_HISTOGRAMS_TAG => { + scratch.histogram.clear(); + message::merge( + wire_type, + &mut scratch.histogram, + &mut buf, + DecodeContext::default(), + ) + .map_err(|mut error| { + error.push("TimeSeries", "histograms"); + error + }) + .context(error::DecodePromRemoteRequestSnafu)?; + rows_remaining = rows_remaining.checked_sub(1).with_context(|| { + error::InvalidPromRemoteRequestSnafu { + msg: "remote write v2 histogram count changed between scans".to_string(), + } + })?; + if let Some(SeriesWriter::Histograms(table_data)) = &mut writer { + if rows_remaining == 0 { + write_native_histogram( + table_data, + &scratch.histogram, + std::mem::take(&mut tags).into_iter(), + )?; + } else { + write_native_histogram( + table_data, + &scratch.histogram, + tags.iter().cloned(), + )?; + } + } + } + TIME_SERIES_EXEMPLARS_TAG => { + scratch.exemplar.clear(); + message::merge( + wire_type, + &mut scratch.exemplar, + &mut buf, + DecodeContext::default(), + ) + .map_err(|mut error| { + error.push("TimeSeries", "exemplars"); + error + }) + .context(error::DecodePromRemoteRequestSnafu)?; + } + TIME_SERIES_METADATA_TAG => { + take_length_delimited(wire_type, &mut buf) + .map_err(|mut error| { + error.push("TimeSeries", "metadata"); + error + }) + .context(error::DecodePromRemoteRequestSnafu)?; + } + _ => skip_field(wire_type, tag, &mut buf, DecodeContext::default()) + .context(error::DecodePromRemoteRequestSnafu)?, + } + } + + ensure!( + rows_remaining == 0, + error::InvalidPromRemoteRequestSnafu { + msg: "remote write v2 row count changed between scans".to_string(), + } + ); + + Ok(()) +} + +fn take_length_delimited<'a>( + wire_type: WireType, + buf: &mut &'a [u8], +) -> std::result::Result<&'a [u8], DecodeError> { + if wire_type != WireType::LengthDelimited { + return Err(DecodeError::new(format!( + "invalid wire type: {wire_type:?} (expected LengthDelimited)" + ))); + } + + let len = decode_varint(buf)?; + let len = + usize::try_from(len).map_err(|_| DecodeError::new("length delimiter exceeds usize"))?; + if len > buf.len() { + return Err(DecodeError::new("buffer underflow")); + } + let (value, remaining) = buf.split_at(len); + *buf = remaining; + Ok(value) +} + +fn checked_total(total: u64, count: usize, name: &str) -> Result { + let count = + u64::try_from(count) + .ok() + .with_context(|| error::InvalidPromRemoteRequestSnafu { + msg: format!("remote write v2 {name} count exceeds u64"), + })?; + total + .checked_add(count) + .with_context(|| error::InvalidPromRemoteRequestSnafu { + msg: format!("remote write v2 {name} count overflows u64"), + }) +} + fn get_or_create_table_data( tables: &mut HashMap>, prom_ctx: PromCtx, @@ -275,46 +588,10 @@ fn get_or_create_table_data( } } -fn write_samples( - table_data: &mut TableData, - mut samples: Vec, - tags: PromTags, -) -> Result<()> { - let Some(last_sample) = samples.pop() else { - return Ok(()); - }; - - for sample in &samples { - write_sample(table_data, sample, tags.iter().cloned())?; - } - - write_sample(table_data, &last_sample, tags.into_iter()) -} - -fn write_sample( - table_data: &mut TableData, - sample: &Sample, - tags: impl Iterator, -) -> Result<()> { - let mut row = table_data.alloc_one_row(); - row_writer::write_ts_to_millis( - table_data, - greptime_timestamp(), - Some(sample.timestamp), - Precision::Millisecond, - &mut row, - )?; - row_writer::write_f64(table_data, greptime_value(), sample.value, &mut row)?; - row_writer::write_tags(table_data, tags, &mut row)?; - table_data.add_row(row); - - Ok(()) -} - -fn write_native_histogram( +fn write_native_histogram<'a>( table_data: &mut TableData, histogram: &Histogram, - tags: impl Iterator, + tags: impl Iterator, ) -> Result<()> { // Persist both int and float families into the logical table schema. Only one // family is populated per row; the other is written as NULL so PromQL can @@ -579,7 +856,7 @@ fn validate_native_histogram_spans( let span_len = spans .iter() .try_fold(0usize, |sum, span| sum.checked_add(span.length as usize)) - .context(error::InvalidPromRemoteRequestSnafu { + .with_context(|| error::InvalidPromRemoteRequestSnafu { msg: format!("remote write v2 native histogram {name} spans overflow"), })?; ensure!( @@ -606,13 +883,13 @@ fn validate_native_histogram_spans( current_index = if span_index == 0 { span.offset } else { - current_index.checked_add(span.offset).context( + current_index.checked_add(span.offset).with_context(|| { error::InvalidPromRemoteRequestSnafu { msg: format!( "remote write v2 native histogram {name} span index overflows i32" ), - }, - )? + } + })? }; for _ in 0..span.length { @@ -701,7 +978,7 @@ fn validate_integer_native_histogram_counts( .try_fold(zero_count, |total, bucket| { total.checked_add(*bucket as u64) }) - .context(error::InvalidPromRemoteRequestSnafu { + .with_context(|| error::InvalidPromRemoteRequestSnafu { msg: "remote write v2 native histogram bucket total overflows u64".to_string(), })?; ensure!( @@ -824,11 +1101,12 @@ fn bucket_counts_from_deltas(deltas: &[i64]) -> Result> { let mut buckets = Vec::with_capacity(deltas.len()); for delta in deltas { - count = count - .checked_add(*delta) - .context(error::InvalidPromRemoteRequestSnafu { - msg: "remote write v2 native histogram bucket count overflows i64".to_string(), - })?; + count = + count + .checked_add(*delta) + .with_context(|| error::InvalidPromRemoteRequestSnafu { + msg: "remote write v2 native histogram bucket count overflows i64".to_string(), + })?; ensure!( count >= 0, error::InvalidPromRemoteRequestSnafu { @@ -841,11 +1119,11 @@ fn bucket_counts_from_deltas(deltas: &[i64]) -> Result> { Ok(buckets) } -fn ensure_no_internal_histogram_labels(tags: &PromTags) -> Result<()> { +fn ensure_no_internal_histogram_labels(tags: &PromTags<'_>) -> Result<()> { // The histogram field column is generated from the protobuf payload. for (name, _) in tags { ensure!( - name != greptime_native_histogram() && name != NATIVE_HISTOGRAM_FIELD, + *name != greptime_native_histogram() && *name != NATIVE_HISTOGRAM_FIELD, error::InvalidPromRemoteRequestSnafu { msg: format!( "remote write v2 label `{name}` conflicts with an internal native histogram label" @@ -858,12 +1136,12 @@ fn ensure_no_internal_histogram_labels(tags: &PromTags) -> Result<()> { } fn resolve_series_labels<'a>( - symbols: &'a [String], - series: &TimeSeries, + symbols: &'a [&str], + labels_refs: &[u32], label_names: &mut HashSet<&'a str>, -) -> Result { +) -> Result> { ensure!( - series.labels_refs.len().is_multiple_of(2), + labels_refs.len().is_multiple_of(2), error::InvalidPromRemoteRequestSnafu { msg: "remote write v2 labels_refs must contain name/value pairs".to_string(), } @@ -871,10 +1149,10 @@ fn resolve_series_labels<'a>( let mut prom_ctx = PromCtx::default(); let mut table_name = None; - let mut tags = Vec::with_capacity(series.labels_refs.len() / 2); + let mut tags = Vec::with_capacity(labels_refs.len() / 2); label_names.clear(); - for pair in series.labels_refs.chunks_exact(2) { + for pair in labels_refs.chunks_exact(2) { let name = symbol_ref(symbols, pair[0], "label name")?; let value = symbol_ref(symbols, pair[1], "label value")?; validate_label(name)?; @@ -893,10 +1171,10 @@ fn resolve_series_labels<'a>( continue; } - tags.push((name.to_string(), value.to_string())); + tags.push((name, value.to_string())); } - let table_name = table_name.context(error::InvalidPromRemoteRequestSnafu { + let table_name = table_name.with_context(|| error::InvalidPromRemoteRequestSnafu { msg: "missing '__name__' label in time-series".to_string(), })?; ensure!( @@ -920,10 +1198,15 @@ fn validate_label(name: &str) -> Result<()> { Ok(()) } -fn symbol_ref<'a>(symbols: &'a [String], idx: u32, field: &str) -> Result<&'a str> { +fn symbol_ref<'a>(symbols: &'a [&str], idx: u32, field: &str) -> Result<&'a str> { + let idx = usize::try_from(idx) + .ok() + .with_context(|| error::InvalidPromRemoteRequestSnafu { + msg: format!("remote write v2 {field} symbol reference exceeds usize"), + })?; symbols - .get(idx as usize) - .map(String::as_str) + .get(idx) + .copied() .with_context(|| error::InvalidPromRemoteRequestSnafu { msg: format!( "remote write v2 {field} symbol reference {idx} is out of range, symbols len: {}", @@ -994,8 +1277,12 @@ pub mod test_util { use api::greptime_proto::io::prometheus::write::v2::{Histogram, Request, Sample, TimeSeries}; use api::v1::RowInsertRequest; use bytes::Bytes; + use prost::Message; + use snafu::ResultExt; - use crate::error::Result; + use crate::error::{self, Result}; + use crate::prom_remote_write::try_decompress; + use crate::prom_store::snappy_compress; pub fn request_with_labels_and_samples( labels: Vec<(&str, &str)>, @@ -1012,13 +1299,42 @@ pub mod test_util { } pub fn decode_request(is_zstd: bool, body: Bytes) -> Result { - super::decode_remote_write_v2_request(is_zstd, body) + let buf = if let Ok(buf) = try_decompress(is_zstd, &body[..]) { + buf + } else { + try_decompress(!is_zstd, &body[..])? + }; + Request::decode(&buf[..]).context(error::DecodePromRemoteRequestSnafu) } pub fn write_requests( request: Request, ) -> Result<(Vec, Vec, u64, u64)> { - let requests = super::into_write_requests(request)?; + let body = Bytes::from(snappy_compress(&request.encode_to_vec())?); + decode_write_requests(false, body, true) + } + + pub fn decode_write_requests( + is_zstd: bool, + body: Bytes, + native_histograms_enabled: bool, + ) -> Result<(Vec, Vec, u64, u64)> { + let requests = super::decode_remote_write_v2(is_zstd, body, native_histograms_enabled)?; + Ok(( + requests.samples.all_req().collect(), + requests.histograms.all_req().collect(), + requests.sample_count, + requests.histogram_count, + )) + } + + pub fn decode_uncompressed_write_requests( + body: &[u8], + native_histograms_enabled: bool, + ) -> Result<(Vec, Vec, u64, u64)> { + let request = + super::BorrowedRequest::decode(body).context(error::DecodePromRemoteRequestSnafu)?; + let requests = super::convert_remote_write_v2(request, native_histograms_enabled)?; Ok(( requests.samples.all_req().collect(), requests.histograms.all_req().collect(), @@ -1109,7 +1425,7 @@ mod tests { let body = Bytes::from(crate::prom_store::snappy_compress(&request.encode_to_vec()).unwrap()); - let decoded = decode_remote_write_v2_request(false, body).unwrap(); + let decoded = test_util::decode_request(false, body.clone()).unwrap(); assert_eq!(decoded.symbols, request.symbols); assert_eq!(decoded.timeseries.len(), 1); @@ -1117,11 +1433,291 @@ mod tests { assert_eq!(decoded.timeseries[0].samples.len(), 1); assert_eq!(decoded.timeseries[0].samples[0].value, 42.0); assert_eq!(decoded.timeseries[0].metadata.as_ref().unwrap().r#type, 1); + assert_eq!( + decode_remote_write_v2(true, body, true) + .unwrap() + .sample_count, + 1 + ); + } + + #[test] + fn test_fused_decoder_accepts_arbitrary_field_order_and_split_labels() { + let mut first_sample = Sample { + value: 42.0, + timestamp: 1000, + start_timestamp: 500, + } + .encode_to_vec(); + first_sample.extend(varint_field(90, 1)); + let second_sample = Sample { + value: 43.0, + timestamp: 2000, + start_timestamp: 1000, + } + .encode_to_vec(); + + let mut series = encoded_message_field(2, &first_sample); + series.extend(packed_u32_field(1, &[1])); + series.extend(varint_field(90, 1)); + series.extend(encoded_message_field(2, &second_sample)); + series.extend(varint_field(1, 2)); + + let mut wire = encoded_message_field(5, &series); + wire.extend(string_field(4, b"")); + wire.extend(varint_field(90, 1)); + wire.extend(string_field(4, METRIC_NAME_LABEL.as_bytes())); + wire.extend(encoded_message_field(5, &packed_u32_field(1, &[99]))); + wire.extend(string_field(4, b"http_requests_total")); + + let requests = decode_wire(&wire, true).unwrap(); + assert_eq!(requests.sample_count, 2); + assert_eq!(requests.histogram_count, 0); + let rows = requests.samples.all_req().next().unwrap().rows.unwrap(); + assert_eq!(rows.rows.len(), 2); + assert_eq!( + rows.schema + .iter() + .map(|column| column.column_name.as_str()) + .collect::>(), + vec![greptime_timestamp(), greptime_value()] + ); + assert_eq!( + rows.rows[0].values[1].value_data, + Some(ValueData::F64Value(42.0)) + ); + assert_eq!( + rows.rows[1].values[1].value_data, + Some(ValueData::F64Value(43.0)) + ); + } + + #[test] + fn test_fused_decoder_accepts_histogram_before_labels_and_unknown_fields() { + let mut histogram = Histogram { + count: Some(Count::CountInt(0)), + zero_count: Some(ZeroCount::ZeroCountInt(0)), + timestamp: 2000, + start_timestamp: 1000, + ..Default::default() + } + .encode_to_vec(); + histogram.extend(varint_field(90, 1)); + + let mut series = encoded_message_field(3, &histogram); + series.extend(packed_u32_field(1, &[1, 2])); + let wire = request_wire(&["", METRIC_NAME_LABEL, "metric"], &[series]); + + let requests = decode_wire(&wire, true).unwrap(); + assert_eq!(requests.histogram_count, 1); + let rows = requests.histograms.all_req().next().unwrap().rows.unwrap(); + assert_eq!( + histogram_field_value(&rows, 0, START_TIMESTAMP_FIELD), + Some(ValueData::TimestampMillisecondValue(1000)) + ); + } + + #[test] + fn test_fused_decoder_rejects_malformed_wire() { + let mut wrong_request_wire = varint_field(4, 0); + let invalid_utf8 = string_field(4, &[0xff]); + + let mut truncated_request = Vec::new(); + prost::encoding::encode_key(5, WireType::LengthDelimited, &mut truncated_request); + prost::encoding::encode_varint(2, &mut truncated_request); + truncated_request.push(0); + + let mut oversized_request = Vec::new(); + prost::encoding::encode_key(4, WireType::LengthDelimited, &mut oversized_request); + prost::encoding::encode_varint(u64::MAX, &mut oversized_request); + + let malformed_varint = vec![0x80; 10]; + + let mut wrong_labels_wire = Vec::new(); + prost::encoding::encode_key(1, WireType::SixtyFourBit, &mut wrong_labels_wire); + wrong_labels_wire.extend([0; 8]); + wrong_labels_wire = request_wire(&[""], &[std::mem::take(&mut wrong_labels_wire)]); + + let wrong_series_wire = request_wire(&[""], &[varint_field(2, 0)]); + + let mut truncated_series = Vec::new(); + prost::encoding::encode_key(2, WireType::LengthDelimited, &mut truncated_series); + prost::encoding::encode_varint(2, &mut truncated_series); + truncated_series.push(0x08); + let truncated_series = request_wire(&[""], &[truncated_series]); + + let mut oversized_series = Vec::new(); + prost::encoding::encode_key(3, WireType::LengthDelimited, &mut oversized_series); + prost::encoding::encode_varint(u64::MAX, &mut oversized_series); + let oversized_series = request_wire(&[""], &[oversized_series]); + + let invalid_sample = series_wire(&[1, 2], 2, &varint_field(1, 1)); + let invalid_sample = request_wire(&["", METRIC_NAME_LABEL, "metric"], &[invalid_sample]); + let invalid_histogram = series_wire(&[1, 2], 3, &varint_field(3, 1)); + let invalid_histogram = + request_wire(&["", METRIC_NAME_LABEL, "metric"], &[invalid_histogram]); + + for (name, wire) in [ + ( + "wrong request wire type", + std::mem::take(&mut wrong_request_wire), + ), + ("invalid utf8", invalid_utf8), + ("truncated request", truncated_request), + ("oversized request length", oversized_request), + ("malformed varint", malformed_varint), + ("wrong labels wire type", wrong_labels_wire), + ("wrong series wire type", wrong_series_wire), + ("truncated series", truncated_series), + ("oversized series length", oversized_series), + ("invalid sample", invalid_sample), + ("invalid histogram", invalid_histogram), + ] { + let error = decode_wire_error(&wire, true, name); + assert!( + matches!(error, error::Error::DecodePromRemoteRequest { .. }), + "{name}: {error}" + ); + } + } + + #[test] + fn test_fused_decoder_rejects_malformed_ignored_messages() { + for tag in [4, 5] { + let mut series = series_wire(&[1, 2], 2, &Sample::default().encode_to_vec()); + series.extend(encoded_message_field(tag, &[0x08])); + let wire = request_wire(&["", METRIC_NAME_LABEL, "metric"], &[series]); + + let error = decode_wire_error(&wire, true, "malformed ignored message"); + assert!(matches!( + error, + error::Error::DecodePromRemoteRequest { .. } + )); + } + } + + #[test] + fn test_fused_decoder_ignores_exemplar_symbol_refs() { + let request = Request { + symbols: vec![ + String::new(), + METRIC_NAME_LABEL.to_string(), + "metric".to_string(), + ], + timeseries: vec![TimeSeries { + labels_refs: vec![1, 2], + samples: vec![Sample::default()], + exemplars: vec![Exemplar { + labels_refs: vec![99], + ..Default::default() + }], + ..Default::default() + }], + }; + + assert_eq!(decode_test_request(request).unwrap().sample_count, 1); + } + + #[test] + fn test_fused_decoder_preserves_empty_request_and_series_behavior() { + let request = Request { + symbols: vec![ + String::new(), + METRIC_NAME_LABEL.to_string(), + "metric".to_string(), + "job".to_string(), + ], + timeseries: vec![ + TimeSeries { + labels_refs: vec![99, 99], + ..Default::default() + }, + TimeSeries { + labels_refs: vec![1], + ..Default::default() + }, + TimeSeries { + labels_refs: vec![3, 2, 3, 2], + ..Default::default() + }, + TimeSeries::default(), + ], + }; + + let requests = decode_test_request(request).unwrap(); + assert_eq!(requests.sample_count, 0); + assert_eq!(requests.histogram_count, 0); + + let requests = decode_test_request(Request { + symbols: vec![String::new()], + timeseries: Vec::new(), + }) + .unwrap(); + assert_eq!(requests.sample_count, 0); + assert!(decode_wire(&[], true).is_err()); + assert!(decode_wire(&[0x0a, 0x00], true).is_err()); + } + + #[test] + fn test_fused_decoder_rejects_duplicate_resolved_label_names() { + let request = Request { + symbols: vec![ + String::new(), + METRIC_NAME_LABEL.to_string(), + "metric".to_string(), + "job".to_string(), + "job".to_string(), + "api".to_string(), + "worker".to_string(), + ], + timeseries: vec![TimeSeries { + labels_refs: vec![1, 2, 3, 5, 4, 6], + samples: vec![Sample::default()], + ..Default::default() + }], + }; + + assert_invalid( + "duplicate resolved labels", + request, + "label name `job` is repeated", + ); + } + + #[test] + fn test_fused_decoder_pins_experimental_error_precedence() { + let histogram = series_wire(&[1, 2], 3, &Histogram::default().encode_to_vec()); + let malformed_sample = series_wire(&[1, 2], 2, &[0x08]); + + let wire = request_wire( + &["", METRIC_NAME_LABEL, "metric"], + &[histogram.clone(), malformed_sample.clone()], + ); + let error = decode_wire_error(&wire, false, "histogram before malformed series"); + assert!(error.to_string().contains("ingestion is experimental")); + + let wire = request_wire( + &["", METRIC_NAME_LABEL, "metric"], + &[malformed_sample, histogram.clone()], + ); + let error = decode_wire_error(&wire, false, "malformed series before histogram"); + assert!(matches!( + error, + error::Error::DecodePromRemoteRequest { .. } + )); + + let missing_name = series_wire(&[3, 4], 2, &Sample::default().encode_to_vec()); + let wire = request_wire( + &["", METRIC_NAME_LABEL, "metric", "job", "api"], + &[missing_name, histogram], + ); + let error = decode_wire_error(&wire, false, "conversion error before histogram"); + assert!(error.to_string().contains("missing '__name__'")); } #[test] fn test_into_context_req_samples() { - let ctx_req = into_write_requests(test_util::request_with_labels_and_samples( + let ctx_req = decode_test_request(test_util::request_with_labels_and_samples( vec![ (METRIC_NAME_LABEL, "http_requests_total"), ("job", "api"), @@ -1183,11 +1779,19 @@ mod tests { rows.rows[1].values[1].value_data, Some(ValueData::F64Value(43.0)) ); + assert_eq!( + rows.rows[1].values[2].value_data, + Some(ValueData::StringValue("api".to_string())) + ); + assert_eq!( + rows.rows[1].values[3].value_data, + Some(ValueData::StringValue("localhost:9090".to_string())) + ); } #[test] fn test_into_context_req_special_labels() { - let ctx_req = into_write_requests(test_util::request_with_labels_and_samples( + let ctx_req = decode_test_request(test_util::request_with_labels_and_samples( vec![ (METRIC_NAME_LABEL, "cpu_usage"), (DATABASE_LABEL, "tenant_a"), @@ -1518,7 +2122,7 @@ mod tests { #[test] fn test_into_context_req_allows_nan_observations_outside_buckets() { - into_write_requests(request_with_histogram(Histogram { + decode_test_request(request_with_histogram(Histogram { count: Some(Count::CountInt(2)), sum: f64::NAN, positive_spans: vec![BucketSpan { @@ -1533,7 +2137,7 @@ mod tests { #[test] fn test_into_context_req_allows_empty_label_values() { - let ctx_req = into_write_requests(test_util::request_with_labels_and_samples( + let ctx_req = decode_test_request(test_util::request_with_labels_and_samples( vec![(METRIC_NAME_LABEL, "metric"), ("job", "")], vec![Sample { value: 1.0, @@ -1649,7 +2253,7 @@ mod tests { }]; histogram.negative_deltas = vec![1]; } - into_write_requests(request_with_histogram(histogram.clone())).unwrap(); + decode_test_request(request_with_histogram(histogram.clone())).unwrap(); let beyond = max_index + 1; if positive { @@ -1694,7 +2298,7 @@ mod tests { ], }; - let ctx_req = into_write_requests(request).unwrap(); + let ctx_req = decode_test_request(request).unwrap(); assert_eq!(ctx_req.sample_count, 1); assert_eq!(ctx_req.histogram_count, 1); @@ -1708,7 +2312,7 @@ mod tests { test_util::request_with_labels_and_samples(vec![(METRIC_NAME_LABEL, "metric")], vec![]); request.timeseries[0].histograms.push(Histogram::default()); - let ctx_req = into_write_requests(request).unwrap(); + let ctx_req = decode_test_request(request).unwrap(); assert_eq!(ctx_req.sample_count, 0); assert_eq!(ctx_req.histogram_count, 1); @@ -1744,7 +2348,7 @@ mod tests { #[test] fn test_into_context_req_preserves_histogram_start_timestamp() { - let ctx_req = into_write_requests(test_util::request_with_labels_and_histograms( + let ctx_req = decode_test_request(test_util::request_with_labels_and_histograms( vec![(METRIC_NAME_LABEL, "metric")], vec![Histogram { timestamp: 2000, @@ -1774,7 +2378,7 @@ mod tests { ); request.timeseries[0].histograms.push(Histogram::default()); - let err = match into_write_requests(request) { + let err = match decode_test_request(request) { Ok(_) => panic!("expected invalid request error"), Err(err) => err, }; @@ -1790,7 +2394,7 @@ mod tests { assert_eq!(greptime_native_histogram(), "custom_native_histogram"); let err = ensure_no_internal_histogram_labels(&vec![( - NATIVE_HISTOGRAM_FIELD.to_string(), + NATIVE_HISTOGRAM_FIELD, "user_value".to_string(), )]) .unwrap_err(); @@ -1837,7 +2441,7 @@ mod tests { ], }; - let ctx_req = into_write_requests(request).unwrap(); + let ctx_req = decode_test_request(request).unwrap(); assert_eq!(ctx_req.histogram_count, 2); let mut inserts = ctx_req.histograms.all_req().collect::>(); @@ -1883,6 +2487,65 @@ mod tests { )); } + fn decode_wire( + wire: &[u8], + native_histograms_enabled: bool, + ) -> Result { + let body = Bytes::from(crate::prom_store::snappy_compress(wire).unwrap()); + decode_remote_write_v2(false, body, native_histograms_enabled) + } + + fn decode_wire_error(wire: &[u8], native_histograms_enabled: bool, name: &str) -> error::Error { + match decode_wire(wire, native_histograms_enabled) { + Ok(_) => panic!("{name}: expected decoder error"), + Err(error) => error, + } + } + + fn request_wire(symbols: &[&str], series: &[Vec]) -> Vec { + let mut wire = Vec::new(); + for symbol in symbols { + wire.extend(string_field(4, symbol.as_bytes())); + } + for series in series { + wire.extend(encoded_message_field(5, series)); + } + wire + } + + fn series_wire(labels_refs: &[u32], leaf_tag: u32, leaf: &[u8]) -> Vec { + let mut series = packed_u32_field(1, labels_refs); + series.extend(encoded_message_field(leaf_tag, leaf)); + series + } + + fn string_field(tag: u32, value: &[u8]) -> Vec { + encoded_message_field(tag, value) + } + + fn encoded_message_field(tag: u32, value: &[u8]) -> Vec { + let mut field = Vec::new(); + prost::encoding::encode_key(tag, WireType::LengthDelimited, &mut field); + prost::encoding::encode_varint(u64::try_from(value.len()).unwrap(), &mut field); + field.extend_from_slice(value); + field + } + + fn varint_field(tag: u32, value: u64) -> Vec { + let mut field = Vec::new(); + prost::encoding::encode_key(tag, WireType::Varint, &mut field); + prost::encoding::encode_varint(value, &mut field); + field + } + + fn packed_u32_field(tag: u32, values: &[u32]) -> Vec { + let mut packed = Vec::new(); + for value in values { + prost::encoding::encode_varint(u64::from(*value), &mut packed); + } + encoded_message_field(tag, &packed) + } + fn request_with_sample(labels: Vec<(&str, &str)>) -> Request { test_util::request_with_labels_and_samples( labels, @@ -1901,8 +2564,21 @@ mod tests { ) } + fn decode_test_request(request: Request) -> Result { + decode_test_request_with_histograms(request, true) + } + + fn decode_test_request_with_histograms( + request: Request, + native_histograms_enabled: bool, + ) -> Result { + let body = + Bytes::from(crate::prom_store::snappy_compress(&request.encode_to_vec()).unwrap()); + decode_remote_write_v2(false, body, native_histograms_enabled) + } + fn assert_invalid(name: &str, request: Request, expected: &str) { - let err = match into_write_requests(request) { + let err = match decode_test_request(request) { Ok(_) => panic!("{name}: expected invalid request error"), Err(err) => err, }; @@ -1993,7 +2669,7 @@ mod tests { >; fn decoded_index(request: Request) -> DecodedIndex { - let encoded = into_write_requests(request) + let encoded = decode_test_request(request) .unwrap() .semantic_index .encode("public") @@ -2045,7 +2721,7 @@ mod tests { assert!(!untyped.contains_key(SEMANTIC_METRIC_TYPE)); assert!(!untyped.contains_key(SEMANTIC_METRIC_METADATA_QUALITY)); - let requests = into_write_requests(metadata_request(vec![( + let requests = decode_test_request(metadata_request(vec![( vec![(METRIC_NAME_LABEL, "untyped_unitless_total")], metadata::MetricType::Unspecified as i32, None, @@ -2054,7 +2730,7 @@ mod tests { assert!(requests.semantic_index.is_empty()); // No metadata at all behaves the same. - let requests = into_write_requests(test_util::request_with_labels_and_samples( + let requests = decode_test_request(test_util::request_with_labels_and_samples( vec![(METRIC_NAME_LABEL, "bare_total")], vec![Sample { value: 1.0, @@ -2123,7 +2799,7 @@ mod tests { None, )]); request.timeseries[0].metadata.as_mut().unwrap().unit_ref = 999; - let err = into_write_requests(request).err().unwrap(); + let err = decode_test_request(request).err().unwrap(); assert!(err.to_string().contains("out of range"), "{err}"); // unit_ref must be validated even when the type is UNSPECIFIED @@ -2134,7 +2810,7 @@ mod tests { None, )]); request.timeseries[0].metadata.as_mut().unwrap().unit_ref = 999; - let err = into_write_requests(request).err().unwrap(); + let err = decode_test_request(request).err().unwrap(); assert!(err.to_string().contains("out of range"), "{err}"); // help_ref is validated although help is never persisted. @@ -2144,7 +2820,7 @@ mod tests { None, )]); request.timeseries[0].metadata.as_mut().unwrap().help_ref = 999; - let err = into_write_requests(request).err().unwrap(); + let err = decode_test_request(request).err().unwrap(); assert!(err.to_string().contains("out of range"), "{err}"); } } diff --git a/src/servers/src/row_writer.rs b/src/servers/src/row_writer.rs index 45afcbca6d..61fd9d77e5 100644 --- a/src/servers/src/row_writer.rs +++ b/src/servers/src/row_writer.rs @@ -92,6 +92,26 @@ impl TableData { Ok(index) } + /// Ensures a default column schema without allocating its name when it already exists. + pub(crate) fn ensure_column_by_name( + &mut self, + name: &str, + datatype: ColumnDataType, + semantic_type: SemanticType, + ) -> Result { + if let Some(index) = self.column_indexes.get(name).copied() { + check_schema(datatype, semantic_type, &self.schema[index])?; + return Ok(index); + } + + self.ensure_column(ColumnSchema { + column_name: name.to_string(), + datatype: datatype as i32, + semantic_type: semantic_type as i32, + ..Default::default() + }) + } + #[allow(dead_code)] pub fn columns(&self) -> &Vec { &self.schema @@ -211,11 +231,14 @@ impl MultiTableData { } /// Write data as tags into the table data. -pub fn write_tags( +pub fn write_tags( table_data: &mut TableData, - tags: impl Iterator, + tags: impl Iterator, one_row: &mut Vec, -) -> Result<()> { +) -> Result<()> +where + K: AsRef + Into, +{ let ktv_iter = tags.map(|(k, v)| (k, ColumnDataType::String, Some(ValueData::StringValue(v)))); write_by_semantic_type(table_data, SemanticType::Tag, ktv_iter, one_row) } @@ -326,12 +349,15 @@ pub(crate) fn write_by_schema( Ok(()) } -fn write_by_semantic_type( +fn write_by_semantic_type( table_data: &mut TableData, semantic_type: SemanticType, - ktv_iter: impl Iterator)>, + ktv_iter: impl Iterator)>, one_row: &mut Vec, -) -> Result<()> { +) -> Result<()> +where + K: AsRef + Into, +{ let TableData { schema, column_indexes, @@ -339,12 +365,13 @@ fn write_by_semantic_type( } = table_data; for (name, datatype, value) in ktv_iter { - let index = column_indexes.get(&name); + let index = column_indexes.get(name.as_ref()).copied(); if let Some(index) = index { - check_schema(datatype, semantic_type, &schema[*index])?; - one_row[*index].value_data = value; + check_schema(datatype, semantic_type, &schema[index])?; + one_row[index].value_data = value; } else { let index = schema.len(); + let name = name.into(); schema.push(ColumnSchema { column_name: name.clone(), datatype: datatype as i32, diff --git a/src/servers/tests/prom_remote_write_v2_test.rs b/src/servers/tests/prom_remote_write_v2_test.rs index 94d260d3fa..75fee046f7 100644 --- a/src/servers/tests/prom_remote_write_v2_test.rs +++ b/src/servers/tests/prom_remote_write_v2_test.rs @@ -77,7 +77,7 @@ fn test_decode_remote_write_v2_native_histogram_dump() { assert_eq!(histogram.timestamp, 1782358160412); let (sample_inserts, histogram_inserts, sample_count, histogram_count) = - remote_write_v2::write_requests(decoded).unwrap(); + remote_write_v2::decode_write_requests(false, Bytes::from_static(BODY), true).unwrap(); assert!(sample_inserts.is_empty()); assert_eq!(sample_count, 0); assert_eq!(histogram_count, 1);