Skip to main content

delta_arrow_reader/reader/
predicate.rs

1//! Logical row predicates and Arrow evaluation.
2
3use std::sync::Arc;
4
5use arrow::{
6    array::{
7        ArrayRef, BinaryArray, BooleanArray, Date32Array, Decimal128Array, FixedSizeBinaryArray,
8        Float32Array, Float64Array, Int8Array, Int16Array, Int32Array, Int64Array,
9        LargeBinaryArray, LargeStringArray, RecordBatch, Scalar as ArrowScalar, StringArray,
10        TimestampMicrosecondArray,
11        types::{Decimal128Type, DecimalType, validate_decimal_precision_and_scale},
12    },
13    compute::{
14        filter_record_batch,
15        kernels::{
16            boolean::{and_kleene, is_not_null, is_null, not, or_kleene},
17            cmp,
18        },
19    },
20    datatypes::{DataType, Schema, TimeUnit},
21    error::ArrowError,
22};
23
24use crate::{DeltaReaderError, error::UnsupportedPredicateSnafu};
25
26/// Comparison operation in a Delta predicate.
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum DeltaComparison {
29    /// Equal.
30    Eq,
31    /// Not equal.
32    NotEq,
33    /// Less than.
34    Lt,
35    /// Less than or equal.
36    LtEq,
37    /// Greater than.
38    Gt,
39    /// Greater than or equal.
40    GtEq,
41}
42
43/// Non-null scalar value in a Delta predicate.
44#[derive(Debug, Clone, PartialEq)]
45pub enum DeltaScalar {
46    /// Boolean value.
47    Boolean(bool),
48    /// Signed 8-bit integer value.
49    Int8(i8),
50    /// Signed 16-bit integer value.
51    Int16(i16),
52    /// Signed 32-bit integer value.
53    Int32(i32),
54    /// Signed 64-bit integer value.
55    Int64(i64),
56    /// 32-bit floating-point value.
57    Float32(f32),
58    /// 64-bit floating-point value.
59    Float64(f64),
60    /// Date value represented as days since the Unix epoch.
61    Date32(i32),
62    /// 128-bit decimal value.
63    Decimal128 {
64        /// Unscaled decimal value.
65        value: i128,
66        /// Decimal precision.
67        precision: u8,
68        /// Decimal scale.
69        scale: i8,
70    },
71    /// UTF-8 string value.
72    Utf8(String),
73    /// Large UTF-8 string value.
74    LargeUtf8(String),
75    /// Binary value.
76    Binary(Vec<u8>),
77    /// Large binary value.
78    LargeBinary(Vec<u8>),
79    /// Fixed-size binary value.
80    FixedSizeBinary {
81        /// Required byte width.
82        size: i32,
83        /// Binary value.
84        value: Vec<u8>,
85    },
86    /// Microsecond timestamp value.
87    TimestampMicrosecond {
88        /// Microseconds since the Unix epoch.
89        value: i64,
90        /// Optional timezone name.
91        timezone: Option<String>,
92    },
93}
94
95/// Query-engine-neutral Delta predicate.
96#[derive(Debug, Clone, PartialEq)]
97pub enum DeltaPredicate {
98    /// Constant Boolean predicate.
99    Constant(bool),
100    /// Compare a column with a non-null scalar value.
101    Compare {
102        /// Unqualified top-level logical column name.
103        column: String,
104        /// Comparison operation.
105        op: DeltaComparison,
106        /// Non-null scalar value.
107        value: DeltaScalar,
108    },
109    /// Test whether a column value is null.
110    IsNull {
111        /// Unqualified top-level logical column name.
112        column: String,
113    },
114    /// Test whether a column value is not null.
115    IsNotNull {
116        /// Unqualified top-level logical column name.
117        column: String,
118    },
119    /// Logical conjunction.
120    And(Vec<DeltaPredicate>),
121    /// Logical disjunction.
122    Or(Vec<DeltaPredicate>),
123    /// Logical negation.
124    Not(Box<DeltaPredicate>),
125}
126
127#[allow(dead_code)]
128pub(crate) fn validate_predicate(
129    predicate: &DeltaPredicate,
130    schema: &Schema,
131) -> Result<(), DeltaReaderError> {
132    match predicate {
133        DeltaPredicate::Constant(_) => Ok(()),
134        DeltaPredicate::Compare { column, value, .. } => {
135            validate_scalar(column_data_type(schema, column)?, value)
136        }
137        DeltaPredicate::IsNull { column } | DeltaPredicate::IsNotNull { column } => {
138            column_data_type(schema, column).map(|_| ())
139        }
140        DeltaPredicate::And(children) | DeltaPredicate::Or(children) => children
141            .iter()
142            .try_for_each(|child| validate_predicate(child, schema)),
143        DeltaPredicate::Not(child) => validate_predicate(child, schema),
144    }
145}
146
147#[allow(dead_code)]
148pub(crate) fn evaluate_predicate(
149    batch: &RecordBatch,
150    predicate: &DeltaPredicate,
151) -> Result<RecordBatch, DeltaReaderError> {
152    predicate_selection(batch, predicate)
153        .and_then(|selection| filter_record_batch(batch, &selection))
154        .map_err(|_| unsupported_predicate("predicate_evaluation"))
155}
156
157pub(crate) fn referenced_columns(predicate: &DeltaPredicate) -> Vec<String> {
158    fn visit(predicate: &DeltaPredicate, columns: &mut Vec<String>) {
159        match predicate {
160            DeltaPredicate::Constant(_) => {}
161            DeltaPredicate::Compare { column, .. }
162            | DeltaPredicate::IsNull { column }
163            | DeltaPredicate::IsNotNull { column } => {
164                if !columns.contains(column) {
165                    columns.push(column.clone());
166                }
167            }
168            DeltaPredicate::And(children) | DeltaPredicate::Or(children) => {
169                for child in children {
170                    visit(child, columns);
171                }
172            }
173            DeltaPredicate::Not(child) => visit(child, columns),
174        }
175    }
176
177    let mut columns = Vec::new();
178    visit(predicate, &mut columns);
179    columns
180}
181
182fn predicate_selection(
183    batch: &RecordBatch,
184    predicate: &DeltaPredicate,
185) -> Result<BooleanArray, ArrowError> {
186    match predicate {
187        DeltaPredicate::Constant(value) => Ok(BooleanArray::from(vec![*value; batch.num_rows()])),
188        DeltaPredicate::Compare { column, op, value } => {
189            let column = batch.column(batch.schema().index_of(column)?);
190            let scalar = ArrowScalar::new(scalar_array(value)?);
191            match op {
192                DeltaComparison::Eq => cmp::eq(column, &scalar),
193                DeltaComparison::NotEq => cmp::neq(column, &scalar),
194                DeltaComparison::Lt => cmp::lt(column, &scalar),
195                DeltaComparison::LtEq => cmp::lt_eq(column, &scalar),
196                DeltaComparison::Gt => cmp::gt(column, &scalar),
197                DeltaComparison::GtEq => cmp::gt_eq(column, &scalar),
198            }
199        }
200        DeltaPredicate::IsNull { column } => {
201            is_null(batch.column(batch.schema().index_of(column)?).as_ref())
202        }
203        DeltaPredicate::IsNotNull { column } => {
204            is_not_null(batch.column(batch.schema().index_of(column)?).as_ref())
205        }
206        DeltaPredicate::And(children) => combine_selections(batch, children, true, and_kleene),
207        DeltaPredicate::Or(children) => combine_selections(batch, children, false, or_kleene),
208        DeltaPredicate::Not(child) => not(&predicate_selection(batch, child)?),
209    }
210}
211
212fn combine_selections(
213    batch: &RecordBatch,
214    predicates: &[DeltaPredicate],
215    identity: bool,
216    combine: fn(&BooleanArray, &BooleanArray) -> Result<BooleanArray, ArrowError>,
217) -> Result<BooleanArray, ArrowError> {
218    let mut predicates = predicates.iter();
219    let Some(first) = predicates.next() else {
220        return Ok(BooleanArray::from(vec![identity; batch.num_rows()]));
221    };
222    let first = predicate_selection(batch, first)?;
223
224    predicates.try_fold(first, |selection, predicate| {
225        combine(&selection, &predicate_selection(batch, predicate)?)
226    })
227}
228
229fn scalar_array(scalar: &DeltaScalar) -> Result<ArrayRef, ArrowError> {
230    let array: ArrayRef = match scalar {
231        DeltaScalar::Boolean(value) => Arc::new(BooleanArray::from(vec![*value])),
232        DeltaScalar::Int8(value) => Arc::new(Int8Array::from(vec![*value])),
233        DeltaScalar::Int16(value) => Arc::new(Int16Array::from(vec![*value])),
234        DeltaScalar::Int32(value) => Arc::new(Int32Array::from(vec![*value])),
235        DeltaScalar::Int64(value) => Arc::new(Int64Array::from(vec![*value])),
236        DeltaScalar::Float32(value) => Arc::new(Float32Array::from(vec![*value])),
237        DeltaScalar::Float64(value) => Arc::new(Float64Array::from(vec![*value])),
238        DeltaScalar::Date32(value) => Arc::new(Date32Array::from(vec![*value])),
239        DeltaScalar::Decimal128 {
240            value,
241            precision,
242            scale,
243        } => Arc::new(
244            Decimal128Array::from(vec![*value]).with_precision_and_scale(*precision, *scale)?,
245        ),
246        DeltaScalar::Utf8(value) => Arc::new(StringArray::from(vec![value.as_str()])),
247        DeltaScalar::LargeUtf8(value) => Arc::new(LargeStringArray::from(vec![value.as_str()])),
248        DeltaScalar::Binary(value) => Arc::new(BinaryArray::from(vec![value.as_slice()])),
249        DeltaScalar::LargeBinary(value) => Arc::new(LargeBinaryArray::from(vec![value.as_slice()])),
250        DeltaScalar::FixedSizeBinary { value, .. } => Arc::new(
251            FixedSizeBinaryArray::try_from_iter(std::iter::once(value.as_slice()))?,
252        ),
253        DeltaScalar::TimestampMicrosecond { value, timezone } => Arc::new(
254            TimestampMicrosecondArray::from(vec![*value]).with_timezone_opt(timezone.clone()),
255        ),
256    };
257    Ok(array)
258}
259
260fn column_data_type<'a>(
261    schema: &'a Schema,
262    column: &str,
263) -> Result<&'a DataType, DeltaReaderError> {
264    if column.is_empty() || column.contains('.') {
265        return Err(unsupported_predicate("invalid_column_reference"));
266    }
267
268    let mut matching_fields = schema
269        .fields()
270        .iter()
271        .filter(|field| field.name() == column);
272    let Some(field) = matching_fields.next() else {
273        return Err(unsupported_predicate("column_not_found"));
274    };
275
276    if matching_fields.next().is_some() {
277        return Err(unsupported_predicate("ambiguous_column"));
278    }
279
280    Ok(field.data_type())
281}
282
283fn validate_scalar(data_type: &DataType, scalar: &DeltaScalar) -> Result<(), DeltaReaderError> {
284    let matches = match scalar {
285        DeltaScalar::Boolean(_) => data_type == &DataType::Boolean,
286        DeltaScalar::Int8(_) => data_type == &DataType::Int8,
287        DeltaScalar::Int16(_) => data_type == &DataType::Int16,
288        DeltaScalar::Int32(_) => data_type == &DataType::Int32,
289        DeltaScalar::Int64(_) => data_type == &DataType::Int64,
290        DeltaScalar::Float32(value) => {
291            if !value.is_finite() {
292                return Err(unsupported_predicate("non_finite_float"));
293            }
294            data_type == &DataType::Float32
295        }
296        DeltaScalar::Float64(value) => {
297            if !value.is_finite() {
298                return Err(unsupported_predicate("non_finite_float"));
299            }
300            data_type == &DataType::Float64
301        }
302        DeltaScalar::Date32(_) => data_type == &DataType::Date32,
303        DeltaScalar::Decimal128 {
304            value,
305            precision,
306            scale,
307        } => {
308            if validate_decimal_precision_and_scale::<Decimal128Type>(*precision, *scale).is_err()
309                || Decimal128Type::validate_decimal_precision(*value, *precision, *scale).is_err()
310            {
311                return Err(unsupported_predicate("invalid_decimal"));
312            }
313            data_type == &DataType::Decimal128(*precision, *scale)
314        }
315        DeltaScalar::Utf8(_) => data_type == &DataType::Utf8,
316        DeltaScalar::LargeUtf8(_) => data_type == &DataType::LargeUtf8,
317        DeltaScalar::Binary(_) => data_type == &DataType::Binary,
318        DeltaScalar::LargeBinary(_) => data_type == &DataType::LargeBinary,
319        DeltaScalar::FixedSizeBinary { size, value } => {
320            if usize::try_from(*size).ok() != Some(value.len()) || *size <= 0 {
321                return Err(unsupported_predicate("invalid_fixed_size_binary"));
322            }
323            data_type == &DataType::FixedSizeBinary(*size)
324        }
325        DeltaScalar::TimestampMicrosecond { timezone, .. } => match timezone {
326            Some(timezone) => {
327                if timezone.is_empty() {
328                    return Err(unsupported_predicate("invalid_timestamp_timezone"));
329                }
330                matches!(
331                    data_type,
332                    DataType::Timestamp(TimeUnit::Microsecond, Some(field_timezone))
333                        if field_timezone.as_ref() == timezone
334                )
335            }
336            None => data_type == &DataType::Timestamp(TimeUnit::Microsecond, None),
337        },
338    };
339
340    if matches {
341        Ok(())
342    } else {
343        Err(unsupported_predicate("scalar_type_mismatch"))
344    }
345}
346
347fn unsupported_predicate(reason: &'static str) -> DeltaReaderError {
348    UnsupportedPredicateSnafu { reason }.build()
349}
350
351#[cfg(test)]
352mod tests {
353    use std::sync::Arc;
354
355    use arrow::{
356        array::{
357            Array, ArrayRef, BinaryArray, BooleanArray, Date32Array, Decimal128Array,
358            FixedSizeBinaryArray, Float32Array, Float64Array, Int8Array, Int16Array, Int32Array,
359            Int64Array, LargeBinaryArray, LargeStringArray, RecordBatch, StringArray,
360            TimestampMicrosecondArray,
361        },
362        datatypes::{
363            DataType, Field, Fields, IntervalUnit, Schema, TimeUnit, UnionFields, UnionMode,
364        },
365    };
366
367    use super::{
368        DeltaComparison, DeltaPredicate, DeltaScalar, evaluate_predicate, predicate_selection,
369        validate_predicate,
370    };
371    use crate::{DeltaReaderError, DeltaReaderPhase};
372
373    fn compare(column: &str, value: DeltaScalar) -> DeltaPredicate {
374        DeltaPredicate::Compare {
375            column: column.into(),
376            op: DeltaComparison::Eq,
377            value,
378        }
379    }
380
381    fn supported_schema() -> Schema {
382        Schema::new(vec![
383            Field::new("boolean", DataType::Boolean, true),
384            Field::new("int8", DataType::Int8, true),
385            Field::new("int16", DataType::Int16, true),
386            Field::new("int32", DataType::Int32, true),
387            Field::new("int64", DataType::Int64, true),
388            Field::new("float32", DataType::Float32, true),
389            Field::new("float64", DataType::Float64, true),
390            Field::new("date32", DataType::Date32, true),
391            Field::new("decimal", DataType::Decimal128(10, 2), true),
392            Field::new("negative_scale", DataType::Decimal128(10, -2), true),
393            Field::new("utf8", DataType::Utf8, true),
394            Field::new("large_utf8", DataType::LargeUtf8, true),
395            Field::new("binary", DataType::Binary, true),
396            Field::new("large_binary", DataType::LargeBinary, true),
397            Field::new("fixed_binary", DataType::FixedSizeBinary(3), true),
398            Field::new(
399                "timestamp",
400                DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
401                true,
402            ),
403            Field::new(
404                "timestamp_ntz",
405                DataType::Timestamp(TimeUnit::Microsecond, None),
406                true,
407            ),
408            Field::new(
409                "struct",
410                DataType::Struct(Fields::from(vec![Field::new(
411                    "nested",
412                    DataType::Int32,
413                    true,
414                )])),
415                true,
416            ),
417        ])
418    }
419
420    fn assert_unsupported(predicate: &DeltaPredicate, schema: &Schema) {
421        let error = validate_predicate(predicate, schema).expect_err("predicate must be rejected");
422        assert_eq!(error.code(), "unsupported_predicate");
423        assert_eq!(error.phase(), DeltaReaderPhase::ScanPlanning);
424    }
425
426    fn selection_values(
427        batch: &RecordBatch,
428        predicate: &DeltaPredicate,
429    ) -> Result<Vec<Option<bool>>, Box<dyn std::error::Error>> {
430        validate_predicate(predicate, batch.schema().as_ref())?;
431        Ok(predicate_selection(batch, predicate)?.iter().collect())
432    }
433
434    fn evaluate_validated(
435        batch: &RecordBatch,
436        predicate: &DeltaPredicate,
437    ) -> Result<RecordBatch, DeltaReaderError> {
438        validate_predicate(predicate, batch.schema().as_ref())?;
439        evaluate_predicate(batch, predicate)
440    }
441
442    #[test]
443    fn accepts_every_exact_scalar_shape_and_predicate_form() -> Result<(), DeltaReaderError> {
444        let schema = supported_schema();
445        let scalars = [
446            ("boolean", DeltaScalar::Boolean(true)),
447            ("int8", DeltaScalar::Int8(i8::MIN)),
448            ("int16", DeltaScalar::Int16(i16::MAX)),
449            ("int32", DeltaScalar::Int32(i32::MIN)),
450            ("int64", DeltaScalar::Int64(i64::MAX)),
451            ("float32", DeltaScalar::Float32(0.0)),
452            ("float32", DeltaScalar::Float32(-0.0)),
453            ("float64", DeltaScalar::Float64(0.0)),
454            ("float64", DeltaScalar::Float64(-0.0)),
455            ("date32", DeltaScalar::Date32(i32::MAX)),
456            (
457                "decimal",
458                DeltaScalar::Decimal128 {
459                    value: 9_999_999_999,
460                    precision: 10,
461                    scale: 2,
462                },
463            ),
464            (
465                "negative_scale",
466                DeltaScalar::Decimal128 {
467                    value: 123,
468                    precision: 10,
469                    scale: -2,
470                },
471            ),
472            ("utf8", DeltaScalar::Utf8(String::new())),
473            ("large_utf8", DeltaScalar::LargeUtf8(String::new())),
474            ("binary", DeltaScalar::Binary(Vec::new())),
475            ("large_binary", DeltaScalar::LargeBinary(Vec::new())),
476            (
477                "fixed_binary",
478                DeltaScalar::FixedSizeBinary {
479                    size: 3,
480                    value: vec![0, 1, 2],
481                },
482            ),
483            (
484                "timestamp",
485                DeltaScalar::TimestampMicrosecond {
486                    value: i64::MAX,
487                    timezone: Some("UTC".into()),
488                },
489            ),
490            (
491                "timestamp_ntz",
492                DeltaScalar::TimestampMicrosecond {
493                    value: i64::MIN,
494                    timezone: None,
495                },
496            ),
497        ];
498        let comparisons = [
499            DeltaComparison::Eq,
500            DeltaComparison::NotEq,
501            DeltaComparison::Lt,
502            DeltaComparison::LtEq,
503            DeltaComparison::Gt,
504            DeltaComparison::GtEq,
505        ];
506
507        for op in comparisons {
508            for (column, value) in &scalars {
509                validate_predicate(
510                    &DeltaPredicate::Compare {
511                        column: (*column).into(),
512                        op,
513                        value: value.clone(),
514                    },
515                    &schema,
516                )?;
517            }
518        }
519
520        validate_predicate(&DeltaPredicate::Constant(true), &schema)?;
521        validate_predicate(
522            &DeltaPredicate::And(vec![
523                DeltaPredicate::IsNull {
524                    column: "struct".into(),
525                },
526                DeltaPredicate::Or(Vec::new()),
527                DeltaPredicate::Not(Box::new(DeltaPredicate::IsNotNull {
528                    column: "int32".into(),
529                })),
530            ]),
531            &schema,
532        )?;
533        validate_predicate(&DeltaPredicate::And(Vec::new()), &schema)?;
534
535        Ok(())
536    }
537
538    #[test]
539    fn rejects_invalid_missing_nested_and_ambiguous_columns_without_disclosure() {
540        let schema = Schema::new(vec![
541            Field::new("id", DataType::Int32, true),
542            Field::new("duplicate", DataType::Int32, true),
543            Field::new("duplicate", DataType::Int32, true),
544        ]);
545
546        for column in ["", "profile.secret", "missing-secret", "duplicate"] {
547            let predicate = compare(column, DeltaScalar::Int32(7));
548            let error = validate_predicate(&predicate, &schema)
549                .expect_err("invalid column must be rejected");
550            let display = error.to_string();
551            assert_eq!(error.code(), "unsupported_predicate");
552            assert_eq!(error.phase(), DeltaReaderPhase::ScanPlanning);
553            assert!(!display.contains("profile"));
554            assert!(!display.contains("missing"));
555            assert!(!display.contains("duplicate"));
556            assert!(!format!("{error:?}").contains("secret"));
557        }
558
559        let hostile_literal = compare("id", DeltaScalar::Utf8("sensitive-literal".into()));
560        let error = validate_predicate(&hostile_literal, &schema)
561            .expect_err("mismatched literal must be rejected");
562        assert!(!error.to_string().contains("sensitive-literal"));
563        assert!(!format!("{error:?}").contains("sensitive-literal"));
564
565        for predicate in [
566            DeltaPredicate::And(vec![
567                DeltaPredicate::Constant(false),
568                compare("missing-secret", DeltaScalar::Int32(7)),
569            ]),
570            DeltaPredicate::Or(vec![
571                DeltaPredicate::Constant(true),
572                compare("missing-secret", DeltaScalar::Int32(7)),
573            ]),
574        ] {
575            assert_unsupported(&predicate, &schema);
576        }
577    }
578
579    #[test]
580    fn rejects_coercion_and_invalid_scalar_values() {
581        let schema = supported_schema();
582        let invalid = [
583            compare("int32", DeltaScalar::Int64(7)),
584            compare("large_utf8", DeltaScalar::Utf8("value".into())),
585            compare("large_binary", DeltaScalar::Binary(vec![1])),
586            compare(
587                "decimal",
588                DeltaScalar::Decimal128 {
589                    value: 1,
590                    precision: 11,
591                    scale: 2,
592                },
593            ),
594            compare(
595                "decimal",
596                DeltaScalar::Decimal128 {
597                    value: 1,
598                    precision: 10,
599                    scale: 3,
600                },
601            ),
602            compare(
603                "decimal",
604                DeltaScalar::Decimal128 {
605                    value: 1,
606                    precision: 0,
607                    scale: 0,
608                },
609            ),
610            compare(
611                "decimal",
612                DeltaScalar::Decimal128 {
613                    value: 10_000_000_000,
614                    precision: 10,
615                    scale: 2,
616                },
617            ),
618            compare("float32", DeltaScalar::Float32(f32::NAN)),
619            compare("float32", DeltaScalar::Float32(f32::INFINITY)),
620            compare("float64", DeltaScalar::Float64(f64::NEG_INFINITY)),
621            compare(
622                "fixed_binary",
623                DeltaScalar::FixedSizeBinary {
624                    size: 0,
625                    value: Vec::new(),
626                },
627            ),
628            compare(
629                "fixed_binary",
630                DeltaScalar::FixedSizeBinary {
631                    size: 3,
632                    value: vec![1, 2],
633                },
634            ),
635            compare(
636                "fixed_binary",
637                DeltaScalar::FixedSizeBinary {
638                    size: 2,
639                    value: vec![1, 2],
640                },
641            ),
642            compare(
643                "timestamp",
644                DeltaScalar::TimestampMicrosecond {
645                    value: 1,
646                    timezone: Some(String::new()),
647                },
648            ),
649            compare(
650                "timestamp",
651                DeltaScalar::TimestampMicrosecond {
652                    value: 1,
653                    timezone: Some("America/Phoenix".into()),
654                },
655            ),
656            compare(
657                "timestamp",
658                DeltaScalar::TimestampMicrosecond {
659                    value: 1,
660                    timezone: None,
661                },
662            ),
663            compare(
664                "timestamp_ntz",
665                DeltaScalar::TimestampMicrosecond {
666                    value: 1,
667                    timezone: Some("UTC".into()),
668                },
669            ),
670        ];
671
672        for predicate in invalid {
673            assert_unsupported(&predicate, &schema);
674        }
675    }
676
677    #[test]
678    fn rejects_every_unsupported_arrow_type() {
679        let item = Arc::new(Field::new("item", DataType::Int32, true));
680        let entries = Arc::new(Field::new(
681            "entries",
682            DataType::Struct(Fields::from(vec![
683                Field::new("key", DataType::Utf8, false),
684                Field::new("value", DataType::Int32, true),
685            ])),
686            false,
687        ));
688        let unsupported = vec![
689            DataType::Null,
690            DataType::UInt8,
691            DataType::UInt16,
692            DataType::UInt32,
693            DataType::UInt64,
694            DataType::Float16,
695            DataType::Date64,
696            DataType::Timestamp(TimeUnit::Second, None),
697            DataType::Timestamp(TimeUnit::Millisecond, None),
698            DataType::Timestamp(TimeUnit::Nanosecond, None),
699            DataType::Time32(TimeUnit::Second),
700            DataType::Time64(TimeUnit::Microsecond),
701            DataType::Duration(TimeUnit::Microsecond),
702            DataType::Interval(IntervalUnit::MonthDayNano),
703            DataType::Decimal32(9, 2),
704            DataType::Decimal64(18, 2),
705            DataType::Decimal256(38, 2),
706            DataType::Utf8View,
707            DataType::BinaryView,
708            DataType::List(Arc::clone(&item)),
709            DataType::Struct(Fields::from(vec![Field::new(
710                "nested",
711                DataType::Int32,
712                true,
713            )])),
714            DataType::Map(entries, false),
715            DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8)),
716            DataType::Union(UnionFields::empty(), UnionMode::Dense),
717        ];
718
719        for data_type in unsupported {
720            let schema = Schema::new(vec![Field::new("value", data_type, true)]);
721            assert_unsupported(&compare("value", DeltaScalar::Int32(7)), &schema);
722        }
723    }
724
725    #[test]
726    fn evaluates_complete_three_valued_truth_tables() -> Result<(), Box<dyn std::error::Error>> {
727        let batch = RecordBatch::try_from_iter([
728            (
729                "a",
730                Arc::new(Int32Array::from(vec![
731                    Some(1),
732                    Some(1),
733                    Some(1),
734                    Some(0),
735                    Some(0),
736                    Some(0),
737                    None,
738                    None,
739                    None,
740                ])) as ArrayRef,
741            ),
742            (
743                "b",
744                Arc::new(Int32Array::from(vec![
745                    Some(1),
746                    Some(0),
747                    None,
748                    Some(1),
749                    Some(0),
750                    None,
751                    Some(1),
752                    Some(0),
753                    None,
754                ])) as ArrayRef,
755            ),
756        ])?;
757        let a = compare("a", DeltaScalar::Int32(1));
758        let b = compare("b", DeltaScalar::Int32(1));
759
760        assert_eq!(
761            selection_values(&batch, &DeltaPredicate::Not(Box::new(a.clone())))?,
762            vec![
763                Some(false),
764                Some(false),
765                Some(false),
766                Some(true),
767                Some(true),
768                Some(true),
769                None,
770                None,
771                None,
772            ]
773        );
774        assert_eq!(
775            selection_values(&batch, &DeltaPredicate::And(vec![a.clone(), b.clone()]))?,
776            vec![
777                Some(true),
778                Some(false),
779                None,
780                Some(false),
781                Some(false),
782                Some(false),
783                None,
784                Some(false),
785                None,
786            ]
787        );
788        assert_eq!(
789            selection_values(&batch, &DeltaPredicate::Or(vec![a, b]))?,
790            vec![
791                Some(true),
792                Some(true),
793                Some(true),
794                Some(true),
795                Some(false),
796                None,
797                Some(true),
798                None,
799                None,
800            ]
801        );
802        assert_eq!(
803            selection_values(&batch, &DeltaPredicate::And(Vec::new()))?,
804            vec![Some(true); 9]
805        );
806        assert_eq!(
807            selection_values(&batch, &DeltaPredicate::Or(Vec::new()))?,
808            vec![Some(false); 9]
809        );
810        assert_eq!(
811            selection_values(&batch, &DeltaPredicate::IsNull { column: "a".into() },)?,
812            vec![
813                Some(false),
814                Some(false),
815                Some(false),
816                Some(false),
817                Some(false),
818                Some(false),
819                Some(true),
820                Some(true),
821                Some(true),
822            ]
823        );
824        assert_eq!(
825            selection_values(&batch, &DeltaPredicate::IsNotNull { column: "a".into() },)?,
826            vec![
827                Some(true),
828                Some(true),
829                Some(true),
830                Some(true),
831                Some(true),
832                Some(true),
833                Some(false),
834                Some(false),
835                Some(false),
836            ]
837        );
838
839        Ok(())
840    }
841
842    #[test]
843    fn evaluates_every_comparison_and_scalar_boundary() -> Result<(), Box<dyn std::error::Error>> {
844        let integer_batch = RecordBatch::try_from_iter([(
845            "value",
846            Arc::new(Int32Array::from(vec![Some(1), Some(2), Some(3), None])) as ArrayRef,
847        )])?;
848        for (op, expected) in [
849            (
850                DeltaComparison::Eq,
851                vec![Some(false), Some(true), Some(false), None],
852            ),
853            (
854                DeltaComparison::NotEq,
855                vec![Some(true), Some(false), Some(true), None],
856            ),
857            (
858                DeltaComparison::Lt,
859                vec![Some(true), Some(false), Some(false), None],
860            ),
861            (
862                DeltaComparison::LtEq,
863                vec![Some(true), Some(true), Some(false), None],
864            ),
865            (
866                DeltaComparison::Gt,
867                vec![Some(false), Some(false), Some(true), None],
868            ),
869            (
870                DeltaComparison::GtEq,
871                vec![Some(false), Some(true), Some(true), None],
872            ),
873        ] {
874            let predicate = DeltaPredicate::Compare {
875                column: "value".into(),
876                op,
877                value: DeltaScalar::Int32(2),
878            };
879            assert_eq!(selection_values(&integer_batch, &predicate)?, expected);
880        }
881
882        let decimal = Decimal128Array::from(vec![123_i128]).with_precision_and_scale(10, 2)?;
883        let negative_decimal =
884            Decimal128Array::from(vec![123_i128]).with_precision_and_scale(10, -2)?;
885        let fixed_binary = FixedSizeBinaryArray::try_from_iter(std::iter::once(b"abc".as_slice()))?;
886        let scalar_batch = RecordBatch::try_from_iter([
887            (
888                "boolean",
889                Arc::new(BooleanArray::from(vec![true])) as ArrayRef,
890            ),
891            ("int8", Arc::new(Int8Array::from(vec![-8])) as ArrayRef),
892            ("int16", Arc::new(Int16Array::from(vec![-16])) as ArrayRef),
893            ("int32", Arc::new(Int32Array::from(vec![-32])) as ArrayRef),
894            ("int64", Arc::new(Int64Array::from(vec![-64])) as ArrayRef),
895            (
896                "float32",
897                Arc::new(Float32Array::from(vec![1.5])) as ArrayRef,
898            ),
899            (
900                "float64",
901                Arc::new(Float64Array::from(vec![-2.5])) as ArrayRef,
902            ),
903            (
904                "date32",
905                Arc::new(Date32Array::from(vec![20_000])) as ArrayRef,
906            ),
907            ("decimal", Arc::new(decimal) as ArrayRef),
908            ("negative_decimal", Arc::new(negative_decimal) as ArrayRef),
909            ("utf8", Arc::new(StringArray::from(vec![""])) as ArrayRef),
910            (
911                "large_utf8",
912                Arc::new(LargeStringArray::from(vec![""])) as ArrayRef,
913            ),
914            (
915                "binary",
916                Arc::new(BinaryArray::from(vec![b"".as_slice()])) as ArrayRef,
917            ),
918            (
919                "large_binary",
920                Arc::new(LargeBinaryArray::from(vec![b"".as_slice()])) as ArrayRef,
921            ),
922            ("fixed_binary", Arc::new(fixed_binary) as ArrayRef),
923            (
924                "timestamp",
925                Arc::new(TimestampMicrosecondArray::from(vec![1_234_567_i64]).with_timezone("UTC"))
926                    as ArrayRef,
927            ),
928            (
929                "timestamp_ntz",
930                Arc::new(TimestampMicrosecondArray::from(vec![1_234_567_i64])) as ArrayRef,
931            ),
932        ])?;
933        let scalars = [
934            ("boolean", DeltaScalar::Boolean(true)),
935            ("int8", DeltaScalar::Int8(-8)),
936            ("int16", DeltaScalar::Int16(-16)),
937            ("int32", DeltaScalar::Int32(-32)),
938            ("int64", DeltaScalar::Int64(-64)),
939            ("float32", DeltaScalar::Float32(1.5)),
940            ("float64", DeltaScalar::Float64(-2.5)),
941            ("date32", DeltaScalar::Date32(20_000)),
942            (
943                "decimal",
944                DeltaScalar::Decimal128 {
945                    value: 123,
946                    precision: 10,
947                    scale: 2,
948                },
949            ),
950            (
951                "negative_decimal",
952                DeltaScalar::Decimal128 {
953                    value: 123,
954                    precision: 10,
955                    scale: -2,
956                },
957            ),
958            ("utf8", DeltaScalar::Utf8(String::new())),
959            ("large_utf8", DeltaScalar::LargeUtf8(String::new())),
960            ("binary", DeltaScalar::Binary(Vec::new())),
961            ("large_binary", DeltaScalar::LargeBinary(Vec::new())),
962            (
963                "fixed_binary",
964                DeltaScalar::FixedSizeBinary {
965                    size: 3,
966                    value: b"abc".to_vec(),
967                },
968            ),
969            (
970                "timestamp",
971                DeltaScalar::TimestampMicrosecond {
972                    value: 1_234_567,
973                    timezone: Some("UTC".into()),
974                },
975            ),
976            (
977                "timestamp_ntz",
978                DeltaScalar::TimestampMicrosecond {
979                    value: 1_234_567,
980                    timezone: None,
981                },
982            ),
983        ];
984
985        for (column, scalar) in scalars {
986            assert_eq!(
987                selection_values(&scalar_batch, &compare(column, scalar))?,
988                vec![Some(true)]
989            );
990        }
991
992        let zero_batch = RecordBatch::try_from_iter([
993            (
994                "float32",
995                Arc::new(Float32Array::from(vec![0.0_f32, -0.0])) as ArrayRef,
996            ),
997            (
998                "float64",
999                Arc::new(Float64Array::from(vec![0.0_f64, -0.0])) as ArrayRef,
1000            ),
1001        ])?;
1002        assert_eq!(
1003            selection_values(&zero_batch, &compare("float32", DeltaScalar::Float32(0.0)))?,
1004            vec![Some(true), Some(false)]
1005        );
1006        assert_eq!(
1007            selection_values(&zero_batch, &compare("float64", DeltaScalar::Float64(-0.0)),)?,
1008            vec![Some(false), Some(true)]
1009        );
1010
1011        Ok(())
1012    }
1013
1014    #[test]
1015    fn filters_sliced_multi_batch_inputs_with_stable_schema_and_order()
1016    -> Result<(), Box<dyn std::error::Error>> {
1017        let full = RecordBatch::try_from_iter([
1018            (
1019                "id",
1020                Arc::new(Int32Array::from(vec![
1021                    Some(99),
1022                    None,
1023                    Some(3),
1024                    Some(1),
1025                    Some(4),
1026                    Some(2),
1027                    Some(88),
1028                ])) as ArrayRef,
1029            ),
1030            (
1031                "label",
1032                Arc::new(StringArray::from(vec!["x", "n", "c", "a", "d", "b", "y"])) as ArrayRef,
1033            ),
1034        ])?;
1035        let batch = full.slice(1, 5);
1036        let predicate = DeltaPredicate::Compare {
1037            column: "id".into(),
1038            op: DeltaComparison::Gt,
1039            value: DeltaScalar::Int32(2),
1040        };
1041        let filtered = evaluate_validated(&batch, &predicate)?;
1042        assert!(Arc::ptr_eq(batch.schema_ref(), filtered.schema_ref()));
1043        assert_eq!(
1044            filtered
1045                .column(0)
1046                .as_any()
1047                .downcast_ref::<Int32Array>()
1048                .ok_or("expected Int32 output")?,
1049            &Int32Array::from(vec![3, 4])
1050        );
1051        assert_eq!(
1052            filtered
1053                .column(1)
1054                .as_any()
1055                .downcast_ref::<StringArray>()
1056                .ok_or("expected Utf8 output")?,
1057            &StringArray::from(vec!["c", "d"])
1058        );
1059
1060        let second = RecordBatch::try_from_iter([(
1061            "id",
1062            Arc::new(Int32Array::from(vec![5, 0])) as ArrayRef,
1063        )])?;
1064        let second_filtered = evaluate_validated(&second, &predicate)?;
1065        assert_eq!(second_filtered.num_rows(), 1);
1066
1067        let no_survivors = evaluate_validated(&batch, &DeltaPredicate::Constant(false))?;
1068        assert_eq!(no_survivors.num_rows(), 0);
1069        assert!(Arc::ptr_eq(batch.schema_ref(), no_survivors.schema_ref()));
1070
1071        let empty = RecordBatch::new_empty(batch.schema());
1072        assert_eq!(evaluate_validated(&empty, &predicate)?.num_rows(), 0);
1073
1074        Ok(())
1075    }
1076
1077    #[test]
1078    fn evaluation_is_stateless_concurrent_and_redacts_failures()
1079    -> Result<(), Box<dyn std::error::Error>> {
1080        let batch = Arc::new(RecordBatch::try_from_iter([(
1081            "id",
1082            Arc::new(Int32Array::from(vec![1, 2, 3])) as ArrayRef,
1083        )])?);
1084        let predicate = Arc::new(DeltaPredicate::Compare {
1085            column: "id".into(),
1086            op: DeltaComparison::GtEq,
1087            value: DeltaScalar::Int32(2),
1088        });
1089        validate_predicate(predicate.as_ref(), batch.schema().as_ref())?;
1090
1091        std::thread::scope(|scope| {
1092            let handles = (0..8)
1093                .map(|_| {
1094                    let batch = Arc::clone(&batch);
1095                    let predicate = Arc::clone(&predicate);
1096                    scope.spawn(move || evaluate_predicate(&batch, &predicate))
1097                })
1098                .collect::<Vec<_>>();
1099
1100            for handle in handles {
1101                let result = handle.join();
1102                assert!(result.is_ok());
1103                assert_eq!(
1104                    result
1105                        .ok()
1106                        .and_then(Result::ok)
1107                        .map(|batch| batch.num_rows()),
1108                    Some(2)
1109                );
1110            }
1111        });
1112
1113        let hostile = compare(
1114            "sensitive-column",
1115            DeltaScalar::Utf8("sensitive-literal".into()),
1116        );
1117        let error = evaluate_predicate(&batch, &hostile)
1118            .expect_err("unexpected evaluation failure must be mapped");
1119        let display = error.to_string();
1120        assert_eq!(error.code(), "unsupported_predicate");
1121        assert!(!display.contains("sensitive-column"));
1122        assert!(!display.contains("sensitive-literal"));
1123        assert!(!format!("{error:?}").contains("sensitive"));
1124
1125        Ok(())
1126    }
1127}