1use std::sync::Arc;
16
17use arc_swap::ArcSwap;
18use common_telemetry::{debug, warn};
19use common_time::Timestamp;
20use common_time::range::TimestampRange;
21use common_time::timestamp::TimeUnit;
22use datafusion::common::ScalarValue;
23use datafusion::physical_optimizer::pruning::PruningPredicateBuilder;
24use datafusion_common::ToDFSchema;
25use datafusion_common::pruning::PruningStatistics;
26use datafusion_common::tree_node::TreeNode;
27use datafusion_expr::expr::{Expr, InList};
28use datafusion_expr::physical_planning_context::PhysicalPlanningContext;
29use datafusion_expr::{Between, BinaryExpr, Operator};
30use datafusion_physical_expr::execution_props::ExecutionProps;
31use datafusion_physical_expr::expressions::{
32 BinaryExpr as PhysicalBinaryExpr, DynamicFilterPhysicalExpr, is_null,
33};
34use datafusion_physical_expr::{PhysicalExpr, create_physical_expr};
35use datatypes::arrow;
36use datatypes::value::scalar_value_to_timestamp;
37use snafu::ResultExt;
38
39use crate::error;
40
41#[cfg(test)]
42mod stats;
43
44macro_rules! return_none_if_utf8 {
47 ($lit: ident) => {
48 if is_string_timestamp_literal($lit) {
49 warn!(
50 "Unexpected ScalarValue::Utf8 in time range predicate: {:?}. Maybe it's an implicit bug, please report it to https://github.com/GreptimeTeam/greptimedb/issues",
51 $lit
52 );
53
54 return None;
56 }
57 };
58}
59
60pub fn is_string_timestamp_literal(scalar: &ScalarValue) -> bool {
61 matches!(
62 scalar,
63 ScalarValue::Utf8(_) | ScalarValue::LargeUtf8(_) | ScalarValue::Utf8View(_)
64 )
65}
66
67#[derive(Debug, Clone, Default)]
69pub struct Predicate {
70 exprs: Arc<Vec<Expr>>,
72 dyn_filters: Arc<ArcSwap<Vec<Arc<DynamicFilterPhysicalExpr>>>>,
76}
77
78impl Predicate {
79 pub fn new(exprs: Vec<Expr>) -> Self {
83 Self {
84 exprs: Arc::new(exprs),
85 dyn_filters: Arc::new(ArcSwap::new(Arc::new(vec![]))),
86 }
87 }
88
89 pub fn with_dyn_filters(
90 exprs: Vec<Expr>,
91 dyn_filters: Vec<Arc<DynamicFilterPhysicalExpr>>,
92 ) -> Self {
93 Self {
94 exprs: Arc::new(exprs),
95 dyn_filters: Arc::new(ArcSwap::new(Arc::new(dyn_filters))),
96 }
97 }
98
99 pub fn is_empty(&self) -> bool {
100 self.exprs.is_empty() && self.dyn_filters.load().is_empty()
101 }
102
103 pub fn add_dyn_filters(&self, dyn_filters: Vec<Arc<DynamicFilterPhysicalExpr>>) {
105 self.dyn_filters.rcu(|existing| {
106 let mut new_filters = existing.as_ref().clone();
107 new_filters.extend(dyn_filters.clone());
108 Arc::new(new_filters)
109 });
110 }
111
112 pub fn clear_dyn_filters(&self) {
114 self.dyn_filters.store(Arc::new(vec![]));
115 }
116
117 pub fn exprs(&self) -> &[Expr] {
119 &self.exprs
120 }
121
122 pub fn dyn_filters(&self) -> Arc<Vec<Arc<DynamicFilterPhysicalExpr>>> {
125 self.dyn_filters.load_full()
126 }
127
128 pub fn dyn_filter_phy_exprs(&self) -> error::Result<Vec<Arc<dyn PhysicalExpr>>> {
131 self.dyn_filters
132 .load()
133 .iter()
134 .map(|e| {
135 e.children()
137 .into_iter()
138 .try_fold(e.current()?, |expr, child| {
139 Ok(Arc::new(PhysicalBinaryExpr::new(
140 expr,
141 Operator::Or,
142 is_null(child.clone())?,
143 )) as Arc<dyn PhysicalExpr>)
144 })
145 })
146 .collect::<Result<Vec<_>, _>>()
147 .context(error::DatafusionSnafu)
148 }
149
150 pub fn to_physical_expr(
152 expr: &Expr,
153 schema: &arrow::datatypes::SchemaRef,
154 ) -> error::Result<Arc<dyn PhysicalExpr>> {
155 let df_schema = schema
156 .clone()
157 .to_dfschema_ref()
158 .context(error::DatafusionSnafu)?;
159
160 let execution_props = &ExecutionProps::new();
164
165 create_physical_expr(
166 expr,
167 df_schema.as_ref(),
168 execution_props,
169 &PhysicalPlanningContext::default(),
170 )
171 .context(error::DatafusionSnafu)
172 }
173
174 pub fn to_physical_exprs(
176 &self,
177 schema: &arrow::datatypes::SchemaRef,
178 ) -> error::Result<Vec<Arc<dyn PhysicalExpr>>> {
179 let dyn_filters = self.dyn_filter_phy_exprs()?;
180
181 Ok(self
182 .exprs
183 .iter()
184 .filter_map(|expr| Self::to_physical_expr(expr, schema).ok())
185 .chain(dyn_filters)
186 .collect::<Vec<_>>())
187 }
188
189 pub fn prune_with_stats<S: PruningStatistics>(
192 &self,
193 stats: &S,
194 schema: &arrow::datatypes::SchemaRef,
195 ) -> Vec<bool> {
196 let mut res = vec![true; stats.num_containers()];
197 let physical_exprs = match self.to_physical_exprs(schema) {
198 Ok(expr) => expr,
199 Err(e) => {
200 warn!(e; "Failed to build physical expr from predicates: {:?}", &self.exprs);
201 return res;
202 }
203 };
204
205 for expr in &physical_exprs {
206 match PruningPredicateBuilder::new()
207 .with_file_schema(schema.clone())
208 .try_build(expr.clone())
209 {
210 Ok(p) => match p.prune(stats) {
211 Ok(r) => {
212 for (curr_val, res) in r.into_iter().zip(res.iter_mut()) {
213 *res &= curr_val
214 }
215 }
216 Err(e) => {
217 warn!(e; "Failed to prune row groups");
218 }
219 },
220 Err(e) => {
221 debug!("Failed to create pruning predicate for expr: {e:?}");
223 }
224 }
225 }
226 res
227 }
228}
229
230pub fn build_time_range_predicate(
235 ts_col_name: &str,
236 ts_col_unit: TimeUnit,
237 filters: &[Expr],
238) -> TimestampRange {
239 let mut res = TimestampRange::min_to_max();
240 for expr in filters {
241 if let Some(range) = extract_time_range_from_expr(ts_col_name, ts_col_unit, expr) {
242 res = res.and(&range);
243 }
244 }
245 res
246}
247
248#[derive(Debug, Clone, PartialEq, Eq)]
255pub enum TimeRangeExtraction {
256 Absent,
258 Extracted(TimestampRange),
261 Unsupported,
264}
265
266pub fn extract_time_range_strict(
269 ts_col_name: &str,
270 ts_col_unit: TimeUnit,
271 filters: &[Expr],
272) -> TimeRangeExtraction {
273 let mut range: Option<TimestampRange> = None;
274 for expr in filters {
275 if !expr
276 .column_refs()
277 .iter()
278 .any(|column| column.name == ts_col_name)
279 {
280 continue;
281 }
282 if contains_disjunction_over_column(expr, ts_col_name) {
287 return TimeRangeExtraction::Unsupported;
288 }
289 match extract_time_range_from_expr(ts_col_name, ts_col_unit, expr) {
294 Some(extracted) => {
295 range = Some(match range {
296 Some(acc) => acc.and(&extracted),
297 None => extracted,
298 });
299 }
300 None => return TimeRangeExtraction::Unsupported,
301 }
302 }
303 match range {
304 Some(range) => TimeRangeExtraction::Extracted(range),
305 None => TimeRangeExtraction::Absent,
306 }
307}
308
309fn contains_disjunction_over_column(expr: &Expr, ts_col_name: &str) -> bool {
310 let is_disjunction = |expr: &Expr| {
311 matches!(
312 expr,
313 Expr::BinaryExpr(BinaryExpr {
314 op: Operator::Or,
315 ..
316 }) | Expr::InList(_)
317 ) && expr
318 .column_refs()
319 .iter()
320 .any(|column| column.name == ts_col_name)
321 };
322 expr.exists(|expr| Ok(is_disjunction(expr))).unwrap_or(true)
323}
324
325pub fn extract_time_range_from_expr(
328 ts_col_name: &str,
329 ts_col_unit: TimeUnit,
330 expr: &Expr,
331) -> Option<TimestampRange> {
332 match expr {
333 Expr::BinaryExpr(BinaryExpr { left, op, right }) => {
334 extract_from_binary_expr(ts_col_name, ts_col_unit, left, op, right)
335 }
336 Expr::Between(Between {
337 expr,
338 negated,
339 low,
340 high,
341 }) => extract_from_between_expr(ts_col_name, ts_col_unit, expr, negated, low, high),
342 Expr::InList(InList {
343 expr,
344 list,
345 negated,
346 }) => extract_from_in_list_expr(ts_col_name, expr, *negated, list),
347 _ => None,
348 }
349}
350
351fn extract_from_binary_expr(
352 ts_col_name: &str,
353 ts_col_unit: TimeUnit,
354 left: &Expr,
355 op: &Operator,
356 right: &Expr,
357) -> Option<TimestampRange> {
358 match op {
359 Operator::Eq => get_timestamp_filter(ts_col_name, left, right)
360 .and_then(|(ts, _)| ts.convert_to(ts_col_unit))
361 .map(TimestampRange::single),
362 Operator::Lt => {
363 let (ts, reverse) = get_timestamp_filter(ts_col_name, left, right)?;
364 if reverse {
365 let ts_val = ts.convert_to(ts_col_unit)?.value();
367 Some(TimestampRange::from_start(Timestamp::new(
368 ts_val + 1,
369 ts_col_unit,
370 )))
371 } else {
372 ts.convert_to_ceil(ts_col_unit)
374 .map(|ts| TimestampRange::until_end(ts, false))
375 }
376 }
377 Operator::LtEq => {
378 let (ts, reverse) = get_timestamp_filter(ts_col_name, left, right)?;
379 if reverse {
380 ts.convert_to_ceil(ts_col_unit)
382 .map(TimestampRange::from_start)
383 } else {
384 ts.convert_to(ts_col_unit)
386 .map(|ts| TimestampRange::until_end(ts, true))
387 }
388 }
389 Operator::Gt => {
390 let (ts, reverse) = get_timestamp_filter(ts_col_name, left, right)?;
391 if reverse {
392 ts.convert_to_ceil(ts_col_unit)
394 .map(|t| TimestampRange::until_end(t, false))
395 } else {
396 let ts_val = ts.convert_to(ts_col_unit)?.value();
398 Some(TimestampRange::from_start(Timestamp::new(
399 ts_val + 1,
400 ts_col_unit,
401 )))
402 }
403 }
404 Operator::GtEq => {
405 let (ts, reverse) = get_timestamp_filter(ts_col_name, left, right)?;
406 if reverse {
407 ts.convert_to(ts_col_unit)
409 .map(|t| TimestampRange::until_end(t, true))
410 } else {
411 ts.convert_to_ceil(ts_col_unit)
413 .map(TimestampRange::from_start)
414 }
415 }
416 Operator::And => {
417 let left = extract_time_range_from_expr(ts_col_name, ts_col_unit, left)
420 .unwrap_or_else(TimestampRange::min_to_max);
421 let right = extract_time_range_from_expr(ts_col_name, ts_col_unit, right)
422 .unwrap_or_else(TimestampRange::min_to_max);
423 Some(left.and(&right))
424 }
425 Operator::Or => {
426 let left = extract_time_range_from_expr(ts_col_name, ts_col_unit, left)?;
427 let right = extract_time_range_from_expr(ts_col_name, ts_col_unit, right)?;
428 Some(left.or(&right))
429 }
430 _ => None,
431 }
432}
433
434fn get_timestamp_filter(ts_col_name: &str, left: &Expr, right: &Expr) -> Option<(Timestamp, bool)> {
435 let (col, lit, reverse) = match (left, right) {
436 (Expr::Column(column), Expr::Literal(scalar, _)) => (column, scalar, false),
437 (Expr::Literal(scalar, _), Expr::Column(column)) => (column, scalar, true),
438 _ => {
439 return None;
440 }
441 };
442 if col.name != ts_col_name {
443 return None;
444 }
445
446 return_none_if_utf8!(lit);
447 scalar_value_to_timestamp(lit, None).map(|t| (t, reverse))
448}
449
450fn extract_from_between_expr(
451 ts_col_name: &str,
452 ts_col_unit: TimeUnit,
453 expr: &Expr,
454 negated: &bool,
455 low: &Expr,
456 high: &Expr,
457) -> Option<TimestampRange> {
458 let Expr::Column(col) = expr else {
459 return None;
460 };
461 if col.name != ts_col_name {
462 return None;
463 }
464
465 if *negated {
466 return None;
467 }
468
469 match (low, high) {
470 (Expr::Literal(low, _), Expr::Literal(high, _)) => {
471 return_none_if_utf8!(low);
472 return_none_if_utf8!(high);
473
474 let low_opt =
475 scalar_value_to_timestamp(low, None).and_then(|ts| ts.convert_to(ts_col_unit));
476 let high_opt = scalar_value_to_timestamp(high, None)
477 .and_then(|ts| ts.convert_to_ceil(ts_col_unit));
478 Some(TimestampRange::new_inclusive(low_opt, high_opt))
479 }
480 _ => None,
481 }
482}
483
484fn extract_from_in_list_expr(
486 ts_col_name: &str,
487 expr: &Expr,
488 negated: bool,
489 list: &[Expr],
490) -> Option<TimestampRange> {
491 if negated {
492 return None;
493 }
494 let Expr::Column(col) = expr else {
495 return None;
496 };
497 if col.name != ts_col_name {
498 return None;
499 }
500
501 if list.is_empty() {
502 return Some(TimestampRange::empty());
503 }
504 let mut init_range = TimestampRange::empty();
505 for expr in list {
506 if let Expr::Literal(scalar, _) = expr {
507 return_none_if_utf8!(scalar);
508 if let Some(timestamp) = scalar_value_to_timestamp(scalar, None) {
509 init_range = init_range.or(&TimestampRange::single(timestamp))
510 } else {
511 return None;
514 }
515 }
516 }
517 Some(init_range)
518}
519
520#[cfg(test)]
521mod tests {
522 use std::sync::Arc;
523
524 use common_test_util::temp_dir::{TempDir, create_temp_dir};
525 use datafusion::parquet::arrow::ArrowWriter;
526 use datafusion_common::{Column, ScalarValue};
527 use datafusion_expr::{BinaryExpr, Literal, Operator, col, lit};
528 use datatypes::arrow::array::Int32Array;
529 use datatypes::arrow::datatypes::{DataType, Field, Schema};
530 use datatypes::arrow::record_batch::RecordBatch;
531 use datatypes::arrow_array::StringArray;
532 use parquet::arrow::ParquetRecordBatchStreamBuilder;
533 use parquet::file::properties::WriterProperties;
534
535 use super::*;
536 use crate::predicate::stats::RowGroupPruningStatistics;
537
538 fn check_build_predicate(expr: Expr, expect: TimestampRange) {
539 assert_eq!(
540 expect,
541 build_time_range_predicate("ts", TimeUnit::Millisecond, &[expr])
542 );
543 }
544
545 #[test]
546 fn test_gt() {
547 check_build_predicate(
549 col("ts").gt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
550 TimestampRange::from_start(Timestamp::new_millisecond(2)),
551 );
552
553 check_build_predicate(
555 lit(ScalarValue::TimestampMillisecond(Some(1), None)).gt(col("ts")),
556 TimestampRange::until_end(Timestamp::new_millisecond(1), false),
557 );
558
559 check_build_predicate(
561 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).gt(col("ts")),
562 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
563 );
564
565 check_build_predicate(
567 col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
568 TimestampRange::from_start(Timestamp::new_millisecond(2)),
569 );
570
571 check_build_predicate(
573 lit(ScalarValue::TimestampSecond(Some(1), None)).gt(col("ts")),
574 TimestampRange::until_end(Timestamp::new_millisecond(1000), false),
575 );
576
577 check_build_predicate(
579 col("ts").gt(lit(ScalarValue::TimestampSecond(Some(1), None))),
580 TimestampRange::from_start(Timestamp::new_millisecond(1001)),
581 );
582 }
583
584 #[test]
585 fn test_gt_eq() {
586 check_build_predicate(
588 col("ts").gt_eq(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
589 TimestampRange::from_start(Timestamp::new_millisecond(1)),
590 );
591
592 check_build_predicate(
594 lit(ScalarValue::TimestampMillisecond(Some(1), None)).gt_eq(col("ts")),
595 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
596 );
597
598 check_build_predicate(
600 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).gt_eq(col("ts")),
601 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
602 );
603
604 check_build_predicate(
606 col("ts").gt_eq(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
607 TimestampRange::from_start(Timestamp::new_millisecond(2)),
608 );
609
610 check_build_predicate(
612 lit(ScalarValue::TimestampSecond(Some(1), None)).gt_eq(col("ts")),
613 TimestampRange::until_end(Timestamp::new_millisecond(1000), true),
614 );
615
616 check_build_predicate(
618 col("ts").gt_eq(lit(ScalarValue::TimestampSecond(Some(1), None))),
619 TimestampRange::from_start(Timestamp::new_millisecond(1000)),
620 );
621 }
622
623 #[test]
624 fn test_lt() {
625 check_build_predicate(
627 col("ts").lt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
628 TimestampRange::until_end(Timestamp::new_millisecond(1), false),
629 );
630
631 check_build_predicate(
633 lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt(col("ts")),
634 TimestampRange::from_start(Timestamp::new_millisecond(2)),
635 );
636
637 check_build_predicate(
639 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt(col("ts")),
640 TimestampRange::from_start(Timestamp::new_millisecond(2)),
641 );
642
643 check_build_predicate(
645 col("ts").lt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
646 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
647 );
648
649 check_build_predicate(
651 lit(ScalarValue::TimestampSecond(Some(1), None)).lt(col("ts")),
652 TimestampRange::from_start(Timestamp::new_millisecond(1001)),
653 );
654
655 check_build_predicate(
657 col("ts").lt(lit(ScalarValue::TimestampSecond(Some(1), None))),
658 TimestampRange::until_end(Timestamp::new_millisecond(1000), false),
659 );
660 }
661
662 #[test]
663 fn test_lt_eq() {
664 check_build_predicate(
666 col("ts").lt_eq(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
667 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
668 );
669
670 check_build_predicate(
672 lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt_eq(col("ts")),
673 TimestampRange::from_start(Timestamp::new_millisecond(1)),
674 );
675
676 check_build_predicate(
678 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt_eq(col("ts")),
679 TimestampRange::from_start(Timestamp::new_millisecond(2)),
680 );
681
682 check_build_predicate(
684 col("ts").lt_eq(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
685 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
686 );
687
688 check_build_predicate(
690 lit(ScalarValue::TimestampSecond(Some(1), None)).lt_eq(col("ts")),
691 TimestampRange::from_start(Timestamp::new_millisecond(1000)),
692 );
693
694 check_build_predicate(
696 col("ts").lt_eq(lit(ScalarValue::TimestampSecond(Some(1), None))),
697 TimestampRange::until_end(Timestamp::new_millisecond(1000), true),
698 );
699 }
700
701 #[test]
702 fn test_extract_time_range_strict() {
703 fn ts_lit(ms: i64) -> Expr {
704 lit(ScalarValue::TimestampMillisecond(Some(ms), None))
705 }
706 let extract =
707 |filters: &[Expr]| extract_time_range_strict("ts", TimeUnit::Millisecond, filters);
708 let range = |start: i64, end: i64| {
709 TimestampRange::new(
710 Timestamp::new_millisecond(start),
711 Timestamp::new_millisecond(end),
712 )
713 .unwrap()
714 };
715
716 assert_eq!(extract(&[]), TimeRangeExtraction::Absent);
718 assert_eq!(
719 extract(&[col("host").eq(lit("a"))]),
720 TimeRangeExtraction::Absent
721 );
722
723 assert_eq!(
725 extract(&[
726 col("ts").gt_eq(ts_lit(1000)),
727 col("ts").lt(ts_lit(2000)),
728 col("host").eq(lit("a")),
729 ]),
730 TimeRangeExtraction::Extracted(range(1000, 2000))
731 );
732
733 assert_eq!(
735 extract(&[col("ts").gt_eq(ts_lit(1000))]),
736 TimeRangeExtraction::Extracted(TimestampRange::from_start(Timestamp::new_millisecond(
737 1000
738 )))
739 );
740 assert_eq!(
741 extract(&[col("ts").lt(ts_lit(2000))]),
742 TimeRangeExtraction::Extracted(TimestampRange::until_end(
743 Timestamp::new_millisecond(2000),
744 false
745 ))
746 );
747
748 assert_eq!(
750 extract(&[col("ts").between(ts_lit(1000), ts_lit(2000))]),
751 TimeRangeExtraction::Extracted(range(1000, 2001))
752 );
753
754 assert_eq!(
756 extract(&[col("ts").eq(ts_lit(1500))]),
757 TimeRangeExtraction::Extracted(TimestampRange::single(Timestamp::new_millisecond(
758 1500
759 )))
760 );
761
762 assert_eq!(
764 extract(&[col("ts").gt_eq(ts_lit(1000)).and(col("ts").lt(col("t2")))]),
765 TimeRangeExtraction::Extracted(TimestampRange::from_start(Timestamp::new_millisecond(
766 1000
767 )))
768 );
769
770 let TimeRangeExtraction::Extracted(empty) =
772 extract(&[col("ts").gt_eq(ts_lit(2000)), col("ts").lt(ts_lit(1000))])
773 else {
774 panic!("expected an extraction");
775 };
776 assert!(empty.is_empty());
777
778 assert_eq!(
781 extract(&[col("ts").gt(ts_lit(1000)).or(col("host").eq(lit("a")))]),
782 TimeRangeExtraction::Unsupported
783 );
784 assert_eq!(
785 extract(&[!col("ts").gt(ts_lit(1000))]),
786 TimeRangeExtraction::Unsupported
787 );
788 assert_eq!(
789 extract(&[col("ts").gt_eq(col("t2"))]),
790 TimeRangeExtraction::Unsupported
791 );
792 }
793
794 async fn gen_test_parquet_file(dir: &TempDir, cnt: usize) -> (String, Arc<Schema>) {
795 let path = dir
796 .path()
797 .join("test-prune.parquet")
798 .to_string_lossy()
799 .to_string();
800
801 let name_field = Field::new("name", DataType::Utf8, true);
802 let count_field = Field::new("cnt", DataType::Int32, true);
803 let schema = Arc::new(Schema::new(vec![name_field, count_field]));
804
805 let file = std::fs::OpenOptions::new()
806 .write(true)
807 .create(true)
808 .truncate(true)
809 .open(path.clone())
810 .unwrap();
811
812 let write_props = WriterProperties::builder()
813 .set_max_row_group_row_count(Some(10))
814 .build();
815 let mut writer = ArrowWriter::try_new(file, schema.clone(), Some(write_props)).unwrap();
816
817 for i in (0..cnt).step_by(10) {
818 let name_array = Arc::new(StringArray::from(
819 (i..(i + 10).min(cnt))
820 .map(|i| i.to_string())
821 .collect::<Vec<_>>(),
822 )) as Arc<_>;
823 let count_array = Arc::new(Int32Array::from(
824 (i..(i + 10).min(cnt)).map(|i| i as i32).collect::<Vec<_>>(),
825 )) as Arc<_>;
826 let rb = RecordBatch::try_new(schema.clone(), vec![name_array, count_array]).unwrap();
827 writer.write(&rb).unwrap();
828 }
829 let _ = writer.close().unwrap();
830 (path, schema)
831 }
832
833 async fn assert_prune(array_cnt: usize, filters: Vec<Expr>, expect: Vec<bool>) {
834 let dir = create_temp_dir("prune_parquet");
835 let (path, arrow_schema) = gen_test_parquet_file(&dir, array_cnt).await;
836 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
837 let arrow_predicate = Predicate::new(filters);
838 let builder = ParquetRecordBatchStreamBuilder::new(
839 tokio::fs::OpenOptions::new()
840 .read(true)
841 .open(path)
842 .await
843 .unwrap(),
844 )
845 .await
846 .unwrap();
847 let metadata = builder.metadata().clone();
848 let row_groups = metadata.row_groups();
849
850 let stats = RowGroupPruningStatistics::new(row_groups, &schema);
851 let res = arrow_predicate.prune_with_stats(&stats, &arrow_schema);
852 assert_eq!(expect, res);
853 }
854
855 #[test]
856 fn test_clear_dyn_filters_preserves_static_predicates() {
857 use datafusion_physical_expr::expressions::lit as physical_lit;
858
859 let static_exprs = vec![col("a").eq(lit(1_i32))];
860 let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(vec![], physical_lit(true)));
861 let predicate =
862 Predicate::with_dyn_filters(static_exprs.clone(), vec![dynamic_filter.clone()]);
863
864 predicate.clear_dyn_filters();
865 dynamic_filter.update(physical_lit(false)).unwrap();
867
868 assert_eq!(predicate.exprs(), static_exprs);
869 assert!(predicate.dyn_filters().is_empty());
870 assert!(predicate.dyn_filter_phy_exprs().unwrap().is_empty());
871 }
872
873 #[tokio::test]
874 async fn test_dynamic_pruning_keeps_null_row_group() {
875 use datafusion_physical_expr::expressions::{
876 Column as PhysicalColumn, lit as physical_lit,
877 };
878
879 let dir = create_temp_dir("dynamic_pruning_nulls");
880 let path = dir.path().join("nullable.parquet");
881 let arrow_schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));
882 let file = std::fs::File::create(&path).unwrap();
883 let mut writer = ArrowWriter::try_new(file, arrow_schema.clone(), None).unwrap();
884 for values in [[None, Some(1)], [Some(1), Some(1)], [None, None]] {
885 let batch = RecordBatch::try_new(
886 arrow_schema.clone(),
887 vec![Arc::new(Int32Array::from(values.to_vec()))],
888 )
889 .unwrap();
890 writer.write(&batch).unwrap();
891 writer.flush().unwrap();
892 }
893 writer.close().unwrap();
894 let builder =
895 ParquetRecordBatchStreamBuilder::new(tokio::fs::File::open(path).await.unwrap())
896 .await
897 .unwrap();
898 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
899 let stats = RowGroupPruningStatistics::new(builder.metadata().row_groups(), &schema);
900 let filter = Arc::new(DynamicFilterPhysicalExpr::new(
901 vec![Arc::new(PhysicalColumn::new("a", 0))],
902 physical_lit(true),
903 ));
904 let predicate = Predicate::with_dyn_filters(vec![], vec![filter.clone()]);
905 assert_eq!(
906 predicate.prune_with_stats(&stats, &arrow_schema),
907 vec![true; 3]
908 );
909 filter
910 .update(
911 Predicate::to_physical_expr(
912 &col("a").gt_eq(lit(10_i32)).and(col("a").lt_eq(lit(10_i32))),
913 &arrow_schema,
914 )
915 .unwrap(),
916 )
917 .unwrap();
918 assert_eq!(
920 predicate.prune_with_stats(&stats, &arrow_schema),
921 vec![true, false, true],
922 );
923 }
924
925 #[tokio::test]
926 async fn test_clear_dyn_filters_restores_static_row_group_pruning() {
927 use datafusion_physical_expr::expressions::{
928 Column as PhysicalColumn, lit as physical_lit,
929 };
930
931 let dir = create_temp_dir("dynamic_pruning_reset");
932 let (path, arrow_schema) = gen_test_parquet_file(&dir, 30).await;
933 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
934 let builder =
935 ParquetRecordBatchStreamBuilder::new(tokio::fs::File::open(path).await.unwrap())
936 .await
937 .unwrap();
938 let stats = RowGroupPruningStatistics::new(builder.metadata().row_groups(), &schema);
939 let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(
940 vec![Arc::new(PhysicalColumn::new("cnt", 1))],
941 physical_lit(true),
942 ));
943 let predicate = Predicate::with_dyn_filters(
944 vec![col("cnt").gt_eq(lit(10_i32))],
945 vec![dynamic_filter.clone()],
946 );
947
948 dynamic_filter
949 .update(
950 Predicate::to_physical_expr(&col("cnt").gt(lit(100_i32)), &arrow_schema).unwrap(),
951 )
952 .unwrap();
953 assert_eq!(
954 predicate.prune_with_stats(&stats, &arrow_schema),
955 vec![false; 3]
956 );
957
958 predicate.clear_dyn_filters();
961 assert_eq!(
962 predicate.prune_with_stats(&stats, &arrow_schema),
963 vec![false, true, true]
964 );
965 }
966
967 fn gen_predicate(max_val: i32, op: Operator) -> Vec<Expr> {
968 vec![datafusion_expr::Expr::BinaryExpr(BinaryExpr {
969 left: Box::new(datafusion_expr::Expr::Column(Column::from_name("cnt"))),
970 op,
971 right: Box::new(max_val.lit()),
972 })]
973 }
974
975 #[tokio::test]
976 async fn test_prune_empty() {
977 assert_prune(3, vec![], vec![true]).await;
978 }
979
980 #[tokio::test]
981 async fn test_prune_all_match() {
982 let p = gen_predicate(3, Operator::Gt);
983 assert_prune(2, p, vec![false]).await;
984 }
985
986 #[tokio::test]
987 async fn test_prune_gt() {
988 let p = gen_predicate(29, Operator::Gt);
989 assert_prune(
990 100,
991 p,
992 vec![
993 false, false, false, true, true, true, true, true, true, true,
994 ],
995 )
996 .await;
997 }
998
999 #[tokio::test]
1000 async fn test_prune_eq_expr() {
1001 let p = gen_predicate(30, Operator::Eq);
1002 assert_prune(40, p, vec![false, false, false, true]).await;
1003 }
1004
1005 #[tokio::test]
1006 async fn test_prune_neq_expr() {
1007 let p = gen_predicate(30, Operator::NotEq);
1008 assert_prune(40, p, vec![true, true, true, true]).await;
1009 }
1010
1011 #[tokio::test]
1012 async fn test_prune_gteq_expr() {
1013 let p = gen_predicate(29, Operator::GtEq);
1014 assert_prune(40, p, vec![false, false, true, true]).await;
1015 }
1016
1017 #[tokio::test]
1018 async fn test_prune_lt_expr() {
1019 let p = gen_predicate(30, Operator::Lt);
1020 assert_prune(40, p, vec![true, true, true, false]).await;
1021 }
1022
1023 #[tokio::test]
1024 async fn test_prune_lteq_expr() {
1025 let p = gen_predicate(30, Operator::LtEq);
1026 assert_prune(40, p, vec![true, true, true, true]).await;
1027 }
1028
1029 #[tokio::test]
1030 async fn test_prune_between_expr() {
1031 let p = gen_predicate(30, Operator::LtEq);
1032 assert_prune(40, p, vec![true, true, true, true]).await;
1033 }
1034
1035 #[tokio::test]
1036 async fn test_or() {
1037 let e = datafusion_expr::Expr::Column(Column::from_name("cnt"))
1039 .gt(30.lit())
1040 .or(datafusion_expr::Expr::Column(Column::from_name("cnt")).lt(20.lit()));
1041 assert_prune(40, vec![e], vec![true, true, false, true]).await;
1042 }
1043
1044 #[tokio::test]
1045 async fn test_to_physical_expr() {
1046 let predicate = Predicate::new(vec![
1047 col("host").eq(lit("host_a")),
1048 col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(123), None))),
1049 ]);
1050
1051 let schema = Arc::new(arrow::datatypes::Schema::new(vec![Field::new(
1052 "host",
1053 arrow::datatypes::DataType::Utf8,
1054 false,
1055 )]));
1056
1057 let predicates = predicate.to_physical_exprs(&schema).unwrap();
1058 assert!(!predicates.is_empty());
1059
1060 let physical_expr = Predicate::to_physical_expr(&col("host").eq(lit("host_a")), &schema);
1061 assert!(physical_expr.is_ok());
1062 }
1063}