use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use arrow_arith::boolean::{and, and_kleene, is_not_null, is_null, not, or, or_kleene};
use arrow_array::cast::AsArray;
use arrow_array::types::{Float32Type, Float64Type};
use arrow_array::{Array, ArrayRef, BooleanArray, Datum as ArrowDatum, RecordBatch, Scalar};
use arrow_buffer::BooleanBuffer;
use arrow_cast::cast::cast;
use arrow_ord::cmp::{eq, gt, gt_eq, lt, lt_eq, neq};
use arrow_schema::{ArrowError, DataType};
use arrow_string::like::starts_with;
use fnv::FnvHashSet;
use parquet::schema::types::SchemaDescriptor;
use crate::arrow::get_arrow_datum;
use crate::error::Result;
use crate::expr::visitors::bound_predicate_visitor::BoundPredicateVisitor;
use crate::expr::{BoundPredicate, BoundReference};
use crate::spec::Datum;
use crate::{Error, ErrorKind};
pub(super) struct CollectFieldIdVisitor {
pub(super) field_ids: HashSet<i32>,
}
impl CollectFieldIdVisitor {
pub(super) fn field_ids(self) -> HashSet<i32> {
self.field_ids
}
}
impl BoundPredicateVisitor for CollectFieldIdVisitor {
type T = ();
fn always_true(&mut self) -> Result<()> {
Ok(())
}
fn always_false(&mut self) -> Result<()> {
Ok(())
}
fn and(&mut self, _lhs: (), _rhs: ()) -> Result<()> {
Ok(())
}
fn or(&mut self, _lhs: (), _rhs: ()) -> Result<()> {
Ok(())
}
fn not(&mut self, _inner: ()) -> Result<()> {
Ok(())
}
fn is_null(&mut self, reference: &BoundReference, _predicate: &BoundPredicate) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn not_null(&mut self, reference: &BoundReference, _predicate: &BoundPredicate) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn is_nan(&mut self, reference: &BoundReference, _predicate: &BoundPredicate) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn not_nan(&mut self, reference: &BoundReference, _predicate: &BoundPredicate) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn less_than(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn less_than_or_eq(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn greater_than(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn greater_than_or_eq(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn eq(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn not_eq(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn starts_with(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn not_starts_with(
&mut self,
reference: &BoundReference,
_literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn r#in(
&mut self,
reference: &BoundReference,
_literals: &FnvHashSet<Datum>,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
fn not_in(
&mut self,
reference: &BoundReference,
_literals: &FnvHashSet<Datum>,
_predicate: &BoundPredicate,
) -> Result<()> {
self.field_ids.insert(reference.field().id);
Ok(())
}
}
pub(super) struct PredicateConverter<'a> {
pub(super) parquet_schema: &'a SchemaDescriptor,
pub(super) column_map: &'a HashMap<i32, usize>,
pub(super) column_indices: &'a Vec<usize>,
}
impl PredicateConverter<'_> {
fn bound_reference(&mut self, reference: &BoundReference) -> Result<Option<usize>> {
if let Some(column_idx) = self.column_map.get(&reference.field().id) {
if self.parquet_schema.get_column_root(*column_idx).is_group() {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Leaf column `{}` in predicates isn't a root column in Parquet schema.",
reference.field().name
),
));
}
let index = self
.column_indices
.iter()
.position(|&idx| idx == *column_idx)
.ok_or(Error::new(
ErrorKind::DataInvalid,
format!(
"Leaf column `{}` in predicates cannot be found in the required column indices.",
reference.field().name
),
))?;
Ok(Some(index))
} else {
Ok(None)
}
}
fn build_always_true(&self) -> Result<Box<PredicateResult>> {
Ok(Box::new(|batch| {
Ok(BooleanArray::from(vec![true; batch.num_rows()]))
}))
}
fn build_always_false(&self) -> Result<Box<PredicateResult>> {
Ok(Box::new(|batch| {
Ok(BooleanArray::from(vec![false; batch.num_rows()]))
}))
}
}
fn project_column(
batch: &RecordBatch,
column_idx: usize,
) -> std::result::Result<ArrayRef, ArrowError> {
let column = batch.column(column_idx);
match column.data_type() {
DataType::Struct(_) => Err(ArrowError::SchemaError(
"Does not support struct column yet.".to_string(),
)),
_ => Ok(column.clone()),
}
}
fn compute_is_nan(array: &ArrayRef) -> std::result::Result<BooleanArray, ArrowError> {
let (is_nan, nulls) = match array.data_type() {
DataType::Float32 => {
let arr = array.as_primitive::<Float32Type>();
(
BooleanBuffer::from_iter(arr.values().iter().map(|v| v.is_nan())),
arr.nulls(),
)
}
DataType::Float64 => {
let arr = array.as_primitive::<Float64Type>();
(
BooleanBuffer::from_iter(arr.values().iter().map(|v| v.is_nan())),
arr.nulls(),
)
}
_ => unreachable!("is_nan is only valid for float types"),
};
let values = match nulls {
Some(nulls) => &is_nan & nulls.inner(),
None => is_nan,
};
Ok(BooleanArray::new(values, None))
}
pub(super) type PredicateResult =
dyn FnMut(RecordBatch) -> std::result::Result<BooleanArray, ArrowError> + Send + 'static;
impl BoundPredicateVisitor for PredicateConverter<'_> {
type T = Box<PredicateResult>;
fn always_true(&mut self) -> Result<Box<PredicateResult>> {
self.build_always_true()
}
fn always_false(&mut self) -> Result<Box<PredicateResult>> {
self.build_always_false()
}
fn and(
&mut self,
mut lhs: Box<PredicateResult>,
mut rhs: Box<PredicateResult>,
) -> Result<Box<PredicateResult>> {
Ok(Box::new(move |batch| {
let left = lhs(batch.clone())?;
let right = rhs(batch)?;
and_kleene(&left, &right)
}))
}
fn or(
&mut self,
mut lhs: Box<PredicateResult>,
mut rhs: Box<PredicateResult>,
) -> Result<Box<PredicateResult>> {
Ok(Box::new(move |batch| {
let left = lhs(batch.clone())?;
let right = rhs(batch)?;
or_kleene(&left, &right)
}))
}
fn not(&mut self, mut inner: Box<PredicateResult>) -> Result<Box<PredicateResult>> {
Ok(Box::new(move |batch| {
let pred_ret = inner(batch)?;
not(&pred_ret)
}))
}
fn is_null(
&mut self,
reference: &BoundReference,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
Ok(Box::new(move |batch| {
let column = project_column(&batch, idx)?;
is_null(&column)
}))
} else {
self.build_always_true()
}
}
fn not_null(
&mut self,
reference: &BoundReference,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
Ok(Box::new(move |batch| {
let column = project_column(&batch, idx)?;
is_not_null(&column)
}))
} else {
self.build_always_false()
}
}
fn is_nan(
&mut self,
reference: &BoundReference,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
Ok(Box::new(move |batch| {
let column = project_column(&batch, idx)?;
compute_is_nan(&column)
}))
} else {
self.build_always_false()
}
}
fn not_nan(
&mut self,
reference: &BoundReference,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
Ok(Box::new(move |batch| {
let column = project_column(&batch, idx)?;
let is_nan = compute_is_nan(&column)?;
not(&is_nan)
}))
} else {
self.build_always_true()
}
}
fn less_than(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
lt(&left, literal.as_ref())
}))
} else {
self.build_always_true()
}
}
fn less_than_or_eq(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
lt_eq(&left, literal.as_ref())
}))
} else {
self.build_always_true()
}
}
fn greater_than(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
gt(&left, literal.as_ref())
}))
} else {
self.build_always_false()
}
}
fn greater_than_or_eq(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
gt_eq(&left, literal.as_ref())
}))
} else {
self.build_always_false()
}
}
fn eq(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
eq(&left, literal.as_ref())
}))
} else {
self.build_always_false()
}
}
fn not_eq(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
neq(&left, literal.as_ref())
}))
} else {
self.build_always_false()
}
}
fn starts_with(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
starts_with(&left, literal.as_ref())
}))
} else {
self.build_always_false()
}
}
fn not_starts_with(
&mut self,
reference: &BoundReference,
literal: &Datum,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literal = get_arrow_datum(literal)?;
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let literal = try_cast_literal(&literal, left.data_type())?;
not(&starts_with(&left, literal.as_ref())?)
}))
} else {
self.build_always_true()
}
}
fn r#in(
&mut self,
reference: &BoundReference,
literals: &FnvHashSet<Datum>,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literals: Vec<_> = literals
.iter()
.map(|lit| get_arrow_datum(lit).unwrap())
.collect();
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let mut acc = BooleanArray::from(vec![false; batch.num_rows()]);
for literal in &literals {
let literal = try_cast_literal(literal, left.data_type())?;
acc = or(&acc, &eq(&left, literal.as_ref())?)?
}
Ok(acc)
}))
} else {
self.build_always_false()
}
}
fn not_in(
&mut self,
reference: &BoundReference,
literals: &FnvHashSet<Datum>,
_predicate: &BoundPredicate,
) -> Result<Box<PredicateResult>> {
if let Some(idx) = self.bound_reference(reference)? {
let literals: Vec<_> = literals
.iter()
.map(|lit| get_arrow_datum(lit).unwrap())
.collect();
Ok(Box::new(move |batch| {
let left = project_column(&batch, idx)?;
let mut acc = BooleanArray::from(vec![true; batch.num_rows()]);
for literal in &literals {
let literal = try_cast_literal(literal, left.data_type())?;
acc = and(&acc, &neq(&left, literal.as_ref())?)?
}
Ok(acc)
}))
} else {
self.build_always_true()
}
}
}
fn try_cast_literal(
literal: &Arc<dyn ArrowDatum + Send + Sync>,
column_type: &DataType,
) -> std::result::Result<Arc<dyn ArrowDatum + Send + Sync>, ArrowError> {
let literal_array = literal.get().0;
if literal_array.data_type() == column_type {
return Ok(Arc::clone(literal));
}
let literal_array = cast(literal_array, column_type)?;
Ok(Arc::new(Scalar::new(literal_array)))
}
#[cfg(test)]
mod tests {
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use arrow_array::{Array, BooleanArray, RecordBatch};
use arrow_schema::{DataType, Field, Schema as ArrowSchema};
use parquet::schema::parser::parse_message_type;
use parquet::schema::types::SchemaDescriptor;
use super::{CollectFieldIdVisitor, PredicateConverter};
use crate::expr::visitors::bound_predicate_visitor::visit;
use crate::expr::{Bind, Predicate, Reference};
use crate::spec::{NestedField, PrimitiveType, Schema, SchemaRef, Type};
fn table_schema_simple() -> SchemaRef {
Arc::new(
Schema::builder()
.with_schema_id(1)
.with_identifier_field_ids(vec![2])
.with_fields(vec![
NestedField::optional(1, "foo", Type::Primitive(PrimitiveType::String)).into(),
NestedField::required(2, "bar", Type::Primitive(PrimitiveType::Int)).into(),
NestedField::optional(3, "baz", Type::Primitive(PrimitiveType::Boolean)).into(),
NestedField::optional(4, "qux", Type::Primitive(PrimitiveType::Float)).into(),
])
.build()
.unwrap(),
)
}
#[test]
fn test_collect_field_id() {
let schema = table_schema_simple();
let expr = Reference::new("qux").is_null();
let bound_expr = expr.bind(schema, true).unwrap();
let mut visitor = CollectFieldIdVisitor {
field_ids: HashSet::default(),
};
visit(&mut visitor, &bound_expr).unwrap();
let mut expected = HashSet::default();
expected.insert(4_i32);
assert_eq!(visitor.field_ids, expected);
}
#[test]
fn test_collect_field_id_with_and() {
let schema = table_schema_simple();
let expr = Reference::new("qux")
.is_null()
.and(Reference::new("baz").is_null());
let bound_expr = expr.bind(schema, true).unwrap();
let mut visitor = CollectFieldIdVisitor {
field_ids: HashSet::default(),
};
visit(&mut visitor, &bound_expr).unwrap();
let mut expected = HashSet::default();
expected.insert(4_i32);
expected.insert(3);
assert_eq!(visitor.field_ids, expected);
}
#[test]
fn test_collect_field_id_with_or() {
let schema = table_schema_simple();
let expr = Reference::new("qux")
.is_null()
.or(Reference::new("baz").is_null());
let bound_expr = expr.bind(schema, true).unwrap();
let mut visitor = CollectFieldIdVisitor {
field_ids: HashSet::default(),
};
visit(&mut visitor, &bound_expr).unwrap();
let mut expected = HashSet::default();
expected.insert(4_i32);
expected.insert(3);
assert_eq!(visitor.field_ids, expected);
}
fn apply_predicate_to_batch(
predicate: Predicate,
schema: SchemaRef,
batch: RecordBatch,
) -> BooleanArray {
let bound = predicate.bind(schema, true).unwrap();
let message_type = "
message schema {
optional float qux = 4;
}
";
let parquet_type = parse_message_type(message_type).expect("parse schema");
let parquet_schema = SchemaDescriptor::new(Arc::new(parquet_type));
let column_map = HashMap::from([(4i32, 0usize)]);
let column_indices = vec![0usize];
let mut converter = PredicateConverter {
parquet_schema: &parquet_schema,
column_map: &column_map,
column_indices: &column_indices,
};
let mut predicate_fn = visit(&mut converter, &bound).unwrap();
predicate_fn(batch).unwrap()
}
#[test]
fn test_predicate_converter_nan() {
use arrow_array::Float32Array;
let schema = table_schema_simple();
let arrow_schema = Arc::new(ArrowSchema::new(vec![Field::new(
"qux",
DataType::Float32,
true,
)]));
let values = vec![Some(1.0f32), Some(f32::NAN), None, Some(0.0f32)];
let batch = RecordBatch::try_new(arrow_schema.clone(), vec![Arc::new(Float32Array::from(
values.clone(),
))])
.unwrap();
let result =
apply_predicate_to_batch(Reference::new("qux").is_nan(), schema.clone(), batch);
assert_eq!(
[
result.value(0),
result.value(1),
result.value(2),
result.value(3)
],
[false, true, false, false]
);
assert!(!result.is_null(2));
let batch =
RecordBatch::try_new(arrow_schema, vec![Arc::new(Float32Array::from(values))]).unwrap();
let result = apply_predicate_to_batch(Reference::new("qux").is_not_nan(), schema, batch);
assert_eq!(
[
result.value(0),
result.value(1),
result.value(2),
result.value(3)
],
[true, false, true, true]
);
assert!(!result.is_null(2));
}
}