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