1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum DeltaComparison {
29 Eq,
31 NotEq,
33 Lt,
35 LtEq,
37 Gt,
39 GtEq,
41}
42
43#[derive(Debug, Clone, PartialEq)]
45pub enum DeltaScalar {
46 Boolean(bool),
48 Int8(i8),
50 Int16(i16),
52 Int32(i32),
54 Int64(i64),
56 Float32(f32),
58 Float64(f64),
60 Date32(i32),
62 Decimal128 {
64 value: i128,
66 precision: u8,
68 scale: i8,
70 },
71 Utf8(String),
73 LargeUtf8(String),
75 Binary(Vec<u8>),
77 LargeBinary(Vec<u8>),
79 FixedSizeBinary {
81 size: i32,
83 value: Vec<u8>,
85 },
86 TimestampMicrosecond {
88 value: i64,
90 timezone: Option<String>,
92 },
93}
94
95#[derive(Debug, Clone, PartialEq)]
97pub enum DeltaPredicate {
98 Constant(bool),
100 Compare {
102 column: String,
104 op: DeltaComparison,
106 value: DeltaScalar,
108 },
109 IsNull {
111 column: String,
113 },
114 IsNotNull {
116 column: String,
118 },
119 And(Vec<DeltaPredicate>),
121 Or(Vec<DeltaPredicate>),
123 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}