1use std::borrow::Borrow;
18use std::collections::HashSet;
19use std::sync::Arc;
20
21use api::v1::SemanticType;
22use datafusion_common::pruning::PruningStatistics;
23use datafusion_common::{Column, ScalarValue};
24use datatypes::arrow::array::{ArrayRef, BooleanArray, UInt64Array};
25use datatypes::data_type::DataType;
26use parquet::file::metadata::RowGroupMetaData;
27use store_api::metadata::RegionMetadataRef;
28use store_api::storage::ColumnId;
29
30use crate::sst::parquet::flat_format::FlatReadFormat;
31use crate::sst::parquet::format::StatValues;
32
33pub(crate) struct RowGroupPruningStats<'a, T> {
35 row_groups: &'a [T],
37 read_format: &'a FlatReadFormat,
39 expected_metadata: Option<RegionMetadataRef>,
44 skip_fields: bool,
46}
47
48impl<'a, T> RowGroupPruningStats<'a, T> {
49 pub(crate) fn new(
51 row_groups: &'a [T],
52 read_format: &'a FlatReadFormat,
53 expected_metadata: Option<RegionMetadataRef>,
54 skip_fields: bool,
55 ) -> Self {
56 Self {
57 row_groups,
58 read_format,
59 expected_metadata,
60 skip_fields,
61 }
62 }
63
64 fn column_id_to_prune(&self, name: &str) -> Option<ColumnId> {
68 let metadata = self
69 .expected_metadata
70 .as_ref()
71 .unwrap_or_else(|| self.read_format.metadata());
72 let col = metadata.column_by_name(name)?;
73
74 if self.skip_fields && col.semantic_type == SemanticType::Field {
76 return None;
77 }
78
79 Some(col.column_id)
80 }
81
82 fn cast_stats_to_expected(&self, column_id: ColumnId, values: ArrayRef) -> Option<ArrayRef> {
93 let Some(expected_metadata) = self.expected_metadata.as_ref() else {
96 return Some(values);
97 };
98 let expected_col = expected_metadata.column_by_id(column_id)?;
99 let file_col = self.read_format.metadata().column_by_id(column_id)?;
100 if expected_col.column_schema.data_type == file_col.column_schema.data_type {
101 return Some(values);
102 }
103 let file_arrow_type = file_col.column_schema.data_type.as_arrow_type();
104 let expected_arrow_type = expected_col.column_schema.data_type.as_arrow_type();
105 let values = if values.data_type() == &file_arrow_type {
106 values
107 } else {
108 datatypes::arrow::compute::cast(&values, &file_arrow_type).ok()?
109 };
110 datatypes::arrow::compute::cast(&values, &expected_arrow_type).ok()
111 }
112
113 fn compat_default_value(&self, column: &str) -> Option<ArrayRef> {
115 let metadata = self.expected_metadata.as_ref()?;
116 let col_metadata = metadata.column_by_name(column)?;
117 col_metadata
118 .column_schema
119 .create_default_vector(self.row_groups.len())
120 .unwrap_or(None)
121 .map(|vector| vector.to_arrow_array())
122 }
123}
124
125impl<T: Borrow<RowGroupMetaData>> RowGroupPruningStats<'_, T> {
126 fn compat_null_count(&self, column: &str) -> Option<ArrayRef> {
128 let metadata = self.expected_metadata.as_ref()?;
129 let col_metadata = metadata.column_by_name(column)?;
130 let value = col_metadata
131 .column_schema
132 .create_default()
133 .unwrap_or(None)?;
134 let values = self.row_groups.iter().map(|meta| {
135 if value.is_null() {
136 u64::try_from(meta.borrow().num_rows()).ok()
137 } else {
138 Some(0)
139 }
140 });
141 Some(Arc::new(UInt64Array::from_iter(values)))
142 }
143}
144
145impl<T: Borrow<RowGroupMetaData>> PruningStatistics for RowGroupPruningStats<'_, T> {
146 fn min_values(&self, column: &Column) -> Option<ArrayRef> {
147 let column_id = self.column_id_to_prune(&column.name)?;
148 match self.read_format.min_values(self.row_groups, column_id) {
149 StatValues::Values(values) => self.cast_stats_to_expected(column_id, values),
150 StatValues::NoColumn => self.compat_default_value(&column.name),
151 StatValues::NoStats => None,
152 }
153 }
154
155 fn max_values(&self, column: &Column) -> Option<ArrayRef> {
156 let column_id = self.column_id_to_prune(&column.name)?;
157 match self.read_format.max_values(self.row_groups, column_id) {
158 StatValues::Values(values) => self.cast_stats_to_expected(column_id, values),
159 StatValues::NoColumn => self.compat_default_value(&column.name),
160 StatValues::NoStats => None,
161 }
162 }
163
164 fn num_containers(&self) -> usize {
165 self.row_groups.len()
166 }
167
168 fn null_counts(&self, column: &Column) -> Option<ArrayRef> {
169 let column_id = self.column_id_to_prune(&column.name)?;
170 match self.read_format.null_counts(self.row_groups, column_id) {
171 StatValues::Values(values) => Some(values),
172 StatValues::NoColumn => self.compat_null_count(&column.name),
173 StatValues::NoStats => None,
174 }
175 }
176
177 fn row_counts(&self, _column: &Column) -> Option<ArrayRef> {
178 None
180 }
181
182 fn contained(&self, _column: &Column, _values: &HashSet<ScalarValue>) -> Option<BooleanArray> {
183 None
185 }
186}
187
188#[cfg(test)]
189mod tests {
190 use std::sync::Arc;
191
192 use datafusion_common::Column;
193 use datatypes::arrow::array::{Array, Int64Array, TimestampMicrosecondArray};
194 use datatypes::arrow::datatypes::TimeUnit;
195 use datatypes::prelude::ConcreteDataType;
196 use parquet::basic::Type as PhysicalType;
197 use parquet::file::metadata::{ColumnChunkMetaData, RowGroupMetaData};
198 use parquet::file::statistics::Statistics;
199 use parquet::schema::types::{SchemaDescriptor, Type};
200 use store_api::codec::PrimaryKeyEncoding;
201 use store_api::metadata::RegionMetadataRef;
202
203 use super::*;
204 use crate::read::read_columns::ReadColumns;
205 use crate::test_util::sst_util::sst_region_metadata_with_encoding;
206
207 fn row_group_with_ts_stats(
211 read_format: &FlatReadFormat,
212 ts_min: i64,
213 ts_max: i64,
214 ) -> RowGroupMetaData {
215 let ts_idx = read_format.arrow_schema().index_of("ts").unwrap();
216 let fields: Vec<Arc<Type>> = read_format
217 .arrow_schema()
218 .fields()
219 .iter()
220 .map(|field| {
221 Arc::new(
222 Type::primitive_type_builder(field.name(), PhysicalType::INT64)
223 .build()
224 .unwrap(),
225 )
226 })
227 .collect();
228 let schema_descr = Arc::new(SchemaDescriptor::new(Arc::new(
229 Type::group_type_builder("schema")
230 .with_fields(fields)
231 .build()
232 .unwrap(),
233 )));
234 let chunks: Vec<_> = (0..schema_descr.num_columns())
235 .map(|i| {
236 let mut builder = ColumnChunkMetaData::builder(schema_descr.column(i));
237 if i == ts_idx {
238 builder = builder.set_statistics(Statistics::int64(
239 Some(ts_min),
240 Some(ts_max),
241 None,
242 Some(0),
243 true,
244 ));
245 }
246 builder.build().unwrap()
247 })
248 .collect();
249 RowGroupMetaData::builder(schema_descr)
250 .set_num_rows(10)
251 .set_total_byte_size(0)
252 .set_column_metadata(chunks)
253 .build()
254 .unwrap()
255 }
256
257 fn expected_metadata_us(file_metadata: &RegionMetadataRef) -> RegionMetadataRef {
260 let mut expected = (**file_metadata).clone();
261 for column in expected.column_metadatas.iter_mut() {
262 if column.column_schema.name == "ts" {
263 column.column_schema.data_type = ConcreteDataType::timestamp_microsecond_datatype();
264 }
265 }
266 Arc::new(expected)
267 }
268
269 fn read_format_for(file_metadata: &RegionMetadataRef) -> FlatReadFormat {
270 FlatReadFormat::new(
271 file_metadata.clone(),
272 ReadColumns::new([0, 1, 2, 3]),
273 None,
274 "test",
275 false,
276 )
277 .unwrap()
278 }
279
280 #[test]
284 fn test_row_group_stats_cast_to_expected_unit() {
285 let file_metadata: RegionMetadataRef =
286 Arc::new(sst_region_metadata_with_encoding(PrimaryKeyEncoding::Dense));
287 let read_format = read_format_for(&file_metadata);
288 let row_group = row_group_with_ts_stats(&read_format, 1_000, 9_000);
289 let column = Column::new_unqualified("ts");
290
291 let groups = [&row_group];
293 let stats = RowGroupPruningStats::new(&groups, &read_format, None, false);
294 let min = stats.min_values(&column).unwrap();
295 let min = min.as_any().downcast_ref::<Int64Array>().unwrap();
296 assert_eq!(1_000, min.value(0));
297
298 let groups = [&row_group];
300 let stats =
301 RowGroupPruningStats::new(&groups, &read_format, Some(file_metadata.clone()), false);
302 let min = stats.min_values(&column).unwrap();
303 let min = min.as_any().downcast_ref::<Int64Array>().unwrap();
304 assert_eq!(1_000, min.value(0));
305
306 let expected = expected_metadata_us(&file_metadata);
308 let groups = [&row_group];
309 let stats = RowGroupPruningStats::new(&groups, &read_format, Some(expected), false);
310 let min = stats.min_values(&column).unwrap();
311 assert_eq!(
312 datatypes::arrow::datatypes::DataType::Timestamp(TimeUnit::Microsecond, None),
313 min.data_type().clone()
314 );
315 let min = min
316 .as_any()
317 .downcast_ref::<TimestampMicrosecondArray>()
318 .unwrap();
319 assert_eq!(1_000_000, min.value(0));
320 let max = stats.max_values(&column).unwrap();
321 let max = max
322 .as_any()
323 .downcast_ref::<TimestampMicrosecondArray>()
324 .unwrap();
325 assert_eq!(9_000_000, max.value(0));
326 }
327}