use crate::Result;
use crate::expr::ExprOperator;
use crate::expr::filter::{Filter, SchemableFilter};
use crate::statistics::{ColumnStatistics, StatisticsContainer};
use arrow_array::{ArrayRef, Datum};
use arrow_ord::cmp;
use arrow_schema::Schema;
use std::collections::HashSet;
#[derive(Debug, Clone)]
pub struct FilePruner {
and_filters: Vec<SchemableFilter>,
}
impl FilePruner {
pub fn new(
and_filters: &[Filter],
table_schema: &Schema,
partition_schema: &Schema,
) -> Result<Self> {
let partition_columns: HashSet<&str> = partition_schema
.fields()
.iter()
.map(|f| f.name().as_str())
.collect();
let and_filters: Vec<SchemableFilter> = and_filters
.iter()
.filter(|filter| !partition_columns.contains(filter.field.as_str()))
.filter_map(|filter| SchemableFilter::try_from((filter.clone(), table_schema)).ok())
.collect();
Ok(FilePruner { and_filters })
}
pub fn empty() -> Self {
FilePruner {
and_filters: Vec::new(),
}
}
pub fn is_empty(&self) -> bool {
self.and_filters.is_empty()
}
pub fn should_include(&self, stats: &StatisticsContainer) -> bool {
if self.and_filters.is_empty() {
return true;
}
for filter in &self.and_filters {
let col_name = filter.field.name();
let Some(col_stats) = stats.columns.get(col_name) else {
continue;
};
if self.can_prune_by_filter(filter, col_stats) {
return false; }
}
true }
fn can_prune_by_filter(&self, filter: &SchemableFilter, col_stats: &ColumnStatistics) -> bool {
if filter.operator.is_multi_value() {
return false;
}
let filter_array = self.extract_filter_array(filter);
let Some(filter_value) = filter_array else {
return false; };
let min = &col_stats.min_value;
let max = &col_stats.max_value;
match filter.operator {
ExprOperator::Eq => {
self.can_prune_eq(&filter_value, min, max)
}
ExprOperator::Ne => {
self.can_prune_ne(&filter_value, min, max)
}
ExprOperator::Lt => {
self.can_prune_lt(&filter_value, min)
}
ExprOperator::Lte => {
self.can_prune_lte(&filter_value, min)
}
ExprOperator::Gt => {
self.can_prune_gt(&filter_value, max)
}
ExprOperator::Gte => {
self.can_prune_gte(&filter_value, max)
}
ExprOperator::In | ExprOperator::NotIn => {
unreachable!("Multi-value operators are short-circuited above")
}
}
}
fn can_prune_eq(
&self,
value: &ArrayRef,
min: &Option<ArrayRef>,
max: &Option<ArrayRef>,
) -> bool {
let Some(min_val) = min else {
return false;
};
let Some(max_val) = max else {
return false;
};
let value_lt_min = cmp::lt(value, min_val).map(|r| r.value(0)).unwrap_or(false);
let value_gt_max = cmp::gt(value, max_val).map(|r| r.value(0)).unwrap_or(false);
value_lt_min || value_gt_max
}
fn can_prune_ne(
&self,
value: &ArrayRef,
min: &Option<ArrayRef>,
max: &Option<ArrayRef>,
) -> bool {
let Some(min_val) = min else {
return false;
};
let Some(max_val) = max else {
return false;
};
let min_eq_max = cmp::eq(min_val, max_val)
.map(|r| r.value(0))
.unwrap_or(false);
let min_eq_value = cmp::eq(min_val, value).map(|r| r.value(0)).unwrap_or(false);
min_eq_max && min_eq_value
}
fn can_prune_lt(&self, value: &ArrayRef, min: &Option<ArrayRef>) -> bool {
let Some(min_val) = min else {
return false;
};
cmp::gt_eq(min_val, value)
.map(|r| r.value(0))
.unwrap_or(false)
}
fn can_prune_lte(&self, value: &ArrayRef, min: &Option<ArrayRef>) -> bool {
let Some(min_val) = min else {
return false;
};
cmp::gt(min_val, value).map(|r| r.value(0)).unwrap_or(false)
}
fn can_prune_gt(&self, value: &ArrayRef, max: &Option<ArrayRef>) -> bool {
let Some(max_val) = max else {
return false;
};
cmp::lt_eq(max_val, value)
.map(|r| r.value(0))
.unwrap_or(false)
}
fn can_prune_gte(&self, value: &ArrayRef, max: &Option<ArrayRef>) -> bool {
let Some(max_val) = max else {
return false;
};
cmp::lt(max_val, value).map(|r| r.value(0)).unwrap_or(false)
}
fn extract_filter_array(&self, filter: &SchemableFilter) -> Option<ArrayRef> {
let (array, is_scalar) = filter.values[0].get();
if array.is_empty() {
return None;
}
if is_scalar || array.len() == 1 {
Some(array.slice(0, 1))
} else {
None
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow_array::{Int64Array, StringArray};
use arrow_schema::{DataType, Field};
use std::sync::Arc;
fn create_test_schema() -> Schema {
Schema::new(vec![
Field::new("id", DataType::Int64, false),
Field::new("name", DataType::Utf8, true),
Field::new("value", DataType::Float64, true),
Field::new("date", DataType::Date32, false),
])
}
fn create_partition_schema() -> Schema {
Schema::new(vec![Field::new("date", DataType::Date32, false)])
}
fn create_stats_with_int_range(col_name: &str, min: i64, max: i64) -> StatisticsContainer {
let mut stats = StatisticsContainer::new(crate::statistics::StatsGranularity::File);
stats.columns.insert(
col_name.to_string(),
ColumnStatistics {
column_name: col_name.to_string(),
data_type: DataType::Int64,
min_value: Some(Arc::new(Int64Array::from(vec![min])) as ArrayRef),
max_value: Some(Arc::new(Int64Array::from(vec![max])) as ArrayRef),
},
);
stats
}
fn create_stats_with_string_range(col_name: &str, min: &str, max: &str) -> StatisticsContainer {
let mut stats = StatisticsContainer::new(crate::statistics::StatsGranularity::File);
stats.columns.insert(
col_name.to_string(),
ColumnStatistics {
column_name: col_name.to_string(),
data_type: DataType::Utf8,
min_value: Some(Arc::new(StringArray::from(vec![min])) as ArrayRef),
max_value: Some(Arc::new(StringArray::from(vec![max])) as ArrayRef),
},
);
stats
}
#[test]
fn test_empty_pruner() {
let pruner = FilePruner::empty();
assert!(pruner.is_empty());
let stats = create_stats_with_int_range("id", 1, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_pruner_excludes_partition_columns() {
let table_schema = create_test_schema();
let partition_schema = create_partition_schema();
let filters = vec![Filter::try_from(("date", "=", "2024-01-01")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
assert!(pruner.is_empty()); }
#[test]
fn test_pruner_keeps_non_partition_columns() {
let table_schema = create_test_schema();
let partition_schema = create_partition_schema();
let filters = vec![Filter::try_from(("id", ">", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
assert!(!pruner.is_empty());
}
#[test]
fn test_eq_filter_prunes_when_value_below_min() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "=", "5")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_eq_filter_prunes_when_value_above_max() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "=", "200")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_eq_filter_includes_when_value_in_range() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_ne_filter_prunes_when_all_equal() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "!=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 50, 50);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_ne_filter_includes_when_range_exists() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "!=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_lt_filter_prunes_when_min_gte_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "<", "10")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_lt_filter_includes_when_min_lt_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "<", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_lte_filter_prunes_when_min_gt_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "<=", "5")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_gt_filter_prunes_when_max_lte_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", ">", "100")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_gt_filter_includes_when_max_gt_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", ">", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_gte_filter_prunes_when_max_lt_value() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", ">=", "150")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
#[test]
fn test_lte_filter_includes_when_value_in_range() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "<=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_gte_filter_includes_when_value_in_range() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", ">=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_string_filter() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("name", "=", "zebra")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_string_range("name", "apple", "banana");
assert!(!pruner.should_include(&stats));
let stats2 = create_stats_with_string_range("name", "apple", "zebra");
assert!(pruner.should_include(&stats2));
}
#[test]
fn test_in_not_in_operators_no_prune() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![
Filter::new(
"id".to_string(),
ExprOperator::In,
vec!["5".to_string(), "10".to_string()],
)
.unwrap(),
];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 50, 100);
assert!(pruner.should_include(&stats));
let filters = vec![
Filter::new(
"id".to_string(),
ExprOperator::NotIn,
vec!["5".to_string(), "10".to_string()],
)
.unwrap(),
];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
assert!(pruner.should_include(&stats));
}
#[test]
fn test_missing_column_stats_includes_file() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("other_column", 1, 10);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_filter_on_column_with_no_stats() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![Filter::try_from(("id", "=", "50")).unwrap()];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let mut stats = StatisticsContainer::new(crate::statistics::StatsGranularity::File);
stats.columns.insert(
"id".to_string(),
ColumnStatistics {
column_name: "id".to_string(),
data_type: DataType::Int64,
min_value: None,
max_value: None,
},
);
assert!(pruner.should_include(&stats));
}
#[test]
fn test_multiple_filters_all_must_pass() {
let table_schema = create_test_schema();
let partition_schema = Schema::empty();
let filters = vec![
Filter::try_from(("id", ">", "0")).unwrap(),
Filter::try_from(("id", "<", "5")).unwrap(),
];
let pruner = FilePruner::new(&filters, &table_schema, &partition_schema).unwrap();
let stats = create_stats_with_int_range("id", 10, 100);
assert!(!pruner.should_include(&stats));
}
}