refactor: simplify bulk memtable stats collection

Signed-off-by: evenyag <realevenyag@gmail.com>
This commit is contained in:
evenyag
2026-07-15 16:21:28 +08:00
parent eaeb0bfbc1
commit b79cf523bc
3 changed files with 120 additions and 32 deletions
+20 -13
View File
@@ -449,8 +449,11 @@ impl Memtable for BulkMemtable {
if bulk_parts.should_compact_unordered_part()
&& let Some(bulk_part) = bulk_parts.unordered_part.to_bulk_part()?
{
let batch_stats =
BatchStats::compute(std::slice::from_ref(&bulk_part.batch), &self.metadata);
let batch_stats = BatchStats::compute(
std::slice::from_ref(&bulk_part.batch),
&self.metadata,
!self.append_mode,
);
bulk_parts.parts.push(BulkPartWrapper {
part: PartToMerge::Bulk {
part: bulk_part,
@@ -462,8 +465,11 @@ impl Memtable for BulkMemtable {
bulk_parts.unordered_part.clear();
}
} else {
let batch_stats =
BatchStats::compute(std::slice::from_ref(&fragment.batch), &self.metadata);
let batch_stats = BatchStats::compute(
std::slice::from_ref(&fragment.batch),
&self.metadata,
!self.append_mode,
);
bulk_parts.parts.push(BulkPartWrapper {
part: PartToMerge::Bulk {
part: fragment,
@@ -515,17 +521,13 @@ impl Memtable for BulkMemtable {
if !bulk_parts.unordered_part.is_empty()
&& let Some(unordered_bulk_part) = bulk_parts.unordered_part.to_bulk_part()?
{
let batch_stats = BatchStats::compute(
std::slice::from_ref(&unordered_bulk_part.batch),
&self.metadata,
);
let part_stats = unordered_bulk_part.to_memtable_stats(&self.metadata);
let range = MemtableRange::new(
Arc::new(MemtableRangeContext::new(
self.id,
Box::new(BulkRangeIterBuilder {
part: unordered_bulk_part,
batch_stats,
batch_stats: None,
context: context.clone(),
sequence,
}),
@@ -550,7 +552,7 @@ impl Memtable for BulkMemtable {
part, batch_stats, ..
} => Box::new(BulkRangeIterBuilder {
part: part.clone(),
batch_stats: batch_stats.clone(),
batch_stats: Some(batch_stats.clone()),
context: context.clone(),
sequence,
}),
@@ -787,7 +789,7 @@ impl BulkMemtable {
/// Iterator builder for bulk range
pub struct BulkRangeIterBuilder {
pub part: BulkPart,
pub(crate) batch_stats: BatchStats,
pub(crate) batch_stats: Option<BatchStats>,
pub context: Arc<BulkIterContext>,
pub sequence: Option<SequenceRange>,
}
@@ -816,7 +818,11 @@ impl IterBuilder for BulkRangeIterBuilder {
_time_range: Option<(Timestamp, Timestamp)>,
metrics: Option<MemScanMetrics>,
) -> Result<BoxedRecordBatchIterator> {
if should_prune_bulk_part(&self.batch_stats, &self.context) {
if self
.batch_stats
.as_ref()
.is_some_and(|stats| should_prune_bulk_part(stats, &self.context))
{
return Ok(Box::new(std::iter::empty()));
}
@@ -1286,6 +1292,7 @@ impl MemtableCompactor {
max_sequence,
estimated_series_count,
metadata,
dedup,
);
common_telemetry::trace!(
@@ -2243,7 +2250,7 @@ mod tests {
/// Helper to create a BulkPartWrapper from a BulkPart.
fn create_bulk_part_wrapper(part: BulkPart) -> BulkPartWrapper {
let metadata = metadata_for_test();
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata);
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata, false);
BulkPartWrapper {
part: PartToMerge::Bulk {
part,
+61 -6
View File
@@ -1270,7 +1270,8 @@ pub(crate) fn should_prune_bulk_part(stats: &BatchStats, context: &BulkIterConte
None => return false,
};
let region_meta = context.read_format().metadata();
let pruning_stats = BatchPruningStats::new(stats, region_meta);
let pruning_stats =
BatchPruningStats::new(stats, region_meta, context.pre_filter_mode().skip_fields());
let mask = predicate.prune_with_stats(&pruning_stats, region_meta.schema.arrow_schema());
!mask.first().copied().unwrap_or(true)
}
@@ -1303,7 +1304,7 @@ impl MultiBulkPart {
pub fn from_bulk_part(part: BulkPart, metadata: &RegionMetadata) -> Self {
let num_rows = part.num_rows();
let series_count = part.estimated_series_count();
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), metadata);
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), metadata, false);
let mut batches = SmallVec::new();
batches.push(part.batch);
@@ -1327,6 +1328,7 @@ impl MultiBulkPart {
/// * `max_sequence` - Maximum sequence number across all batches
/// * `series_count` - Number of series in the batches
/// * `metadata` - Region metadata for computing batch statistics
/// * `skip_fields` - Whether to skip field column statistics
///
/// # Panics
/// Panics if batches is empty.
@@ -1337,11 +1339,12 @@ impl MultiBulkPart {
max_sequence: SequenceNumber,
series_count: usize,
metadata: &RegionMetadata,
skip_fields: bool,
) -> Self {
assert!(!batches.is_empty(), "batches must not be empty");
let total_rows = batches.iter().map(|b| b.num_rows()).sum();
let batch_stats = BatchStats::compute(&batches, metadata);
let batch_stats = BatchStats::compute(&batches, metadata, skip_fields);
Self {
batches: SmallVec::from_vec(batches),
@@ -1428,7 +1431,11 @@ impl MultiBulkPart {
fn prune_batches(&self, context: &BulkIterContextRef) -> Vec<RecordBatch> {
if let Some(predicate) = &context.predicate {
let region_meta = context.read_format().metadata();
let pruning_stats = BatchPruningStats::new(&self.batch_stats, region_meta);
let pruning_stats = BatchPruningStats::new(
&self.batch_stats,
region_meta,
context.pre_filter_mode().skip_fields(),
);
let mask =
predicate.prune_with_stats(&pruning_stats, region_meta.schema.arrow_schema());
self.batches
@@ -2598,6 +2605,7 @@ mod tests {
max_seq,
groups.len(),
&metadata,
false,
);
(multi, metadata)
}
@@ -2723,7 +2731,19 @@ mod tests {
false,
)
.unwrap();
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata);
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata, false);
assert!(should_prune_bulk_part(&batch_stats, &context));
// The stored k0 column is dictionary encoded. Stats use the dictionary values array.
let context = BulkIterContext::new(
metadata.clone(),
None,
Some(Predicate::new(vec![
datafusion_expr::col("k0").eq(datafusion_expr::lit("missing")),
])),
false,
)
.unwrap();
assert!(should_prune_bulk_part(&batch_stats, &context));
let context = BulkIterContext::new(
@@ -2738,6 +2758,41 @@ mod tests {
assert!(should_prune_bulk_part(&batch_stats, &context));
}
#[test]
fn test_bulk_part_stats_skip_fields() {
let input = [MutationInput {
k0: "a",
k1: 10,
timestamps: &[1, 2],
v1: &[Some(10.0), Some(11.0)],
sequence: 0,
}];
let metadata = metadata_for_test();
let part = build_converted_bulk_part(&input);
let stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata, true);
let tag_context = BulkIterContext::new(
metadata.clone(),
None,
Some(Predicate::new(vec![
datafusion_expr::col("k1").eq(datafusion_expr::lit(999u32)),
])),
false,
)
.unwrap();
assert!(should_prune_bulk_part(&stats, &tag_context));
let field_context = BulkIterContext::new(
metadata,
None,
Some(Predicate::new(vec![
datafusion_expr::col("v1").gt(datafusion_expr::lit(100.0f64)),
])),
false,
)
.unwrap();
assert!(!should_prune_bulk_part(&stats, &field_context));
}
#[test]
fn test_bulk_part_minmax_all_null_field_keeps_batch() {
let metadata = metadata_for_test();
@@ -2755,7 +2810,7 @@ mod tests {
false,
)
.unwrap();
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata);
let batch_stats = BatchStats::compute(std::slice::from_ref(&part.batch), &metadata, false);
assert!(!should_prune_bulk_part(&batch_stats, &context));
}
}
+39 -13
View File
@@ -17,6 +17,7 @@
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use api::v1::SemanticType;
use datafusion_common::pruning::PruningStatistics;
use datafusion_common::{Column, ScalarValue};
use datatypes::arrow;
@@ -34,12 +35,9 @@ use datatypes::arrow::datatypes::{
i256,
};
use datatypes::data_type::{ConcreteDataType, DataType};
use snafu::ResultExt;
use store_api::metadata::{ColumnMetadata, RegionMetadata, RegionMetadataRef};
use store_api::storage::ColumnId;
use crate::error::{ComputeArrowSnafu, Result};
type ScalarPair = (ScalarValue, ScalarValue);
/// Per-batch min/max statistics for columns in a [`MultiBulkPart`](crate::memtable::bulk::part::MultiBulkPart).
@@ -62,10 +60,14 @@ impl BatchStats {
pub(crate) fn compute(
batches: &[common_recordbatch::DfRecordBatch],
metadata: &RegionMetadata,
skip_fields: bool,
) -> Self {
let mut columns = HashMap::with_capacity(metadata.column_metadatas.len());
for column in &metadata.column_metadatas {
if skip_fields && column.semantic_type == SemanticType::Field {
continue;
}
let Some(stats) = compute_column_stats(batches, column) else {
continue;
};
@@ -120,7 +122,13 @@ fn compute_column_stats(
let mut maxes = Vec::with_capacity(batches.len());
for batch in batches {
let Some((min, max)) = min_max_scalar(batch.column(column_idx), &arrow_type).ok()? else {
let array = batch.column(column_idx);
if array.data_type() != &arrow_type
&& !matches!(array.data_type(), ArrowDataType::Dictionary(_, value_type) if value_type.as_ref() == &arrow_type)
{
return None;
}
let Some((min, max)) = min_max_scalar(array, &arrow_type) else {
mins.push(null_scalar.clone());
maxes.push(null_scalar.clone());
continue;
@@ -176,16 +184,19 @@ fn is_supported_arrow_type(arrow_type: &ArrowDataType) -> bool {
/// Returns exact min/max scalars for a supported Arrow array.
///
/// Empty/all-null arrays return `Ok(None)`. Unsupported arrays return `Ok(None)`.
fn min_max_scalar(array: &ArrayRef, logical_type: &ArrowDataType) -> Result<Option<ScalarPair>> {
/// Empty/all-null and unsupported arrays return `None`.
fn min_max_scalar(array: &ArrayRef, logical_type: &ArrowDataType) -> Option<ScalarPair> {
if array.is_empty() || array.null_count() == array.len() {
return Ok(None);
return None;
}
if let ArrowDataType::Dictionary(_, value_type) = array.data_type() {
let decoded =
arrow::compute::cast(array.as_ref(), value_type).context(ComputeArrowSnafu)?;
return min_max_scalar(&decoded, value_type);
let values = array.as_any_dictionary().values();
let logical_type = match logical_type {
ArrowDataType::Dictionary(_, logical_value_type) => logical_value_type,
_ => value_type,
};
return min_max_scalar(&values, logical_type);
}
let stats = match logical_type {
@@ -274,7 +285,7 @@ fn min_max_scalar(array: &ArrayRef, logical_type: &ArrowDataType) -> Result<Opti
_ => None,
};
Ok(stats)
stats
}
trait ScalarValueFromPrimitive: ArrowPrimitiveType {
@@ -433,23 +444,38 @@ fn fixed_size_binary_min_max(array: &FixedSizeBinaryArray) -> Option<ScalarPair>
pub(crate) struct BatchPruningStats<'a> {
stats: &'a BatchStats,
metadata: &'a RegionMetadataRef,
skip_fields: bool,
}
impl<'a> BatchPruningStats<'a> {
/// Creates a new [`BatchPruningStats`].
pub(crate) fn new(stats: &'a BatchStats, metadata: &'a RegionMetadataRef) -> Self {
Self { stats, metadata }
pub(crate) fn new(
stats: &'a BatchStats,
metadata: &'a RegionMetadataRef,
skip_fields: bool,
) -> Self {
Self {
stats,
metadata,
skip_fields,
}
}
}
impl PruningStatistics for BatchPruningStats<'_> {
fn min_values(&self, column: &Column) -> Option<ArrayRef> {
let col = self.metadata.column_by_name(&column.name)?;
if self.skip_fields && col.semantic_type == SemanticType::Field {
return None;
}
self.stats.min_values(col.column_id)
}
fn max_values(&self, column: &Column) -> Option<ArrayRef> {
let col = self.metadata.column_by_name(&column.name)?;
if self.skip_fields && col.semantic_type == SemanticType::Field {
return None;
}
self.stats.max_values(col.column_id)
}