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 let timestamp = scalar_value_to_timestamp(scalar, None)?;
511 init_range = init_range.or(&TimestampRange::single(timestamp));
512 }
513 }
514 Some(init_range)
515}
516
517#[cfg(test)]
518mod tests {
519 use std::sync::Arc;
520
521 use common_test_util::temp_dir::{TempDir, create_temp_dir};
522 use datafusion::parquet::arrow::ArrowWriter;
523 use datafusion_common::{Column, ScalarValue};
524 use datafusion_expr::{BinaryExpr, Literal, Operator, col, lit};
525 use datatypes::arrow::array::Int32Array;
526 use datatypes::arrow::datatypes::{DataType, Field, Schema};
527 use datatypes::arrow::record_batch::RecordBatch;
528 use datatypes::arrow_array::StringArray;
529 use parquet::arrow::ParquetRecordBatchStreamBuilder;
530 use parquet::file::properties::WriterProperties;
531
532 use super::*;
533 use crate::predicate::stats::RowGroupPruningStatistics;
534
535 fn check_build_predicate(expr: Expr, expect: TimestampRange) {
536 assert_eq!(
537 expect,
538 build_time_range_predicate("ts", TimeUnit::Millisecond, &[expr])
539 );
540 }
541
542 #[test]
543 fn test_gt() {
544 check_build_predicate(
546 col("ts").gt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
547 TimestampRange::from_start(Timestamp::new_millisecond(2)),
548 );
549
550 check_build_predicate(
552 lit(ScalarValue::TimestampMillisecond(Some(1), None)).gt(col("ts")),
553 TimestampRange::until_end(Timestamp::new_millisecond(1), false),
554 );
555
556 check_build_predicate(
558 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).gt(col("ts")),
559 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
560 );
561
562 check_build_predicate(
564 col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
565 TimestampRange::from_start(Timestamp::new_millisecond(2)),
566 );
567
568 check_build_predicate(
570 lit(ScalarValue::TimestampSecond(Some(1), None)).gt(col("ts")),
571 TimestampRange::until_end(Timestamp::new_millisecond(1000), false),
572 );
573
574 check_build_predicate(
576 col("ts").gt(lit(ScalarValue::TimestampSecond(Some(1), None))),
577 TimestampRange::from_start(Timestamp::new_millisecond(1001)),
578 );
579 }
580
581 #[test]
582 fn test_gt_eq() {
583 check_build_predicate(
585 col("ts").gt_eq(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
586 TimestampRange::from_start(Timestamp::new_millisecond(1)),
587 );
588
589 check_build_predicate(
591 lit(ScalarValue::TimestampMillisecond(Some(1), None)).gt_eq(col("ts")),
592 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
593 );
594
595 check_build_predicate(
597 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).gt_eq(col("ts")),
598 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
599 );
600
601 check_build_predicate(
603 col("ts").gt_eq(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
604 TimestampRange::from_start(Timestamp::new_millisecond(2)),
605 );
606
607 check_build_predicate(
609 lit(ScalarValue::TimestampSecond(Some(1), None)).gt_eq(col("ts")),
610 TimestampRange::until_end(Timestamp::new_millisecond(1000), true),
611 );
612
613 check_build_predicate(
615 col("ts").gt_eq(lit(ScalarValue::TimestampSecond(Some(1), None))),
616 TimestampRange::from_start(Timestamp::new_millisecond(1000)),
617 );
618 }
619
620 #[test]
621 fn test_lt() {
622 check_build_predicate(
624 col("ts").lt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
625 TimestampRange::until_end(Timestamp::new_millisecond(1), false),
626 );
627
628 check_build_predicate(
630 lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt(col("ts")),
631 TimestampRange::from_start(Timestamp::new_millisecond(2)),
632 );
633
634 check_build_predicate(
636 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt(col("ts")),
637 TimestampRange::from_start(Timestamp::new_millisecond(2)),
638 );
639
640 check_build_predicate(
642 col("ts").lt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
643 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
644 );
645
646 check_build_predicate(
648 lit(ScalarValue::TimestampSecond(Some(1), None)).lt(col("ts")),
649 TimestampRange::from_start(Timestamp::new_millisecond(1001)),
650 );
651
652 check_build_predicate(
654 col("ts").lt(lit(ScalarValue::TimestampSecond(Some(1), None))),
655 TimestampRange::until_end(Timestamp::new_millisecond(1000), false),
656 );
657 }
658
659 #[test]
660 fn test_lt_eq() {
661 check_build_predicate(
663 col("ts").lt_eq(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
664 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
665 );
666
667 check_build_predicate(
669 lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt_eq(col("ts")),
670 TimestampRange::from_start(Timestamp::new_millisecond(1)),
671 );
672
673 check_build_predicate(
675 lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt_eq(col("ts")),
676 TimestampRange::from_start(Timestamp::new_millisecond(2)),
677 );
678
679 check_build_predicate(
681 col("ts").lt_eq(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
682 TimestampRange::until_end(Timestamp::new_millisecond(1), true),
683 );
684
685 check_build_predicate(
687 lit(ScalarValue::TimestampSecond(Some(1), None)).lt_eq(col("ts")),
688 TimestampRange::from_start(Timestamp::new_millisecond(1000)),
689 );
690
691 check_build_predicate(
693 col("ts").lt_eq(lit(ScalarValue::TimestampSecond(Some(1), None))),
694 TimestampRange::until_end(Timestamp::new_millisecond(1000), true),
695 );
696 }
697
698 #[test]
699 fn test_extract_time_range_strict() {
700 fn ts_lit(ms: i64) -> Expr {
701 lit(ScalarValue::TimestampMillisecond(Some(ms), None))
702 }
703 let extract =
704 |filters: &[Expr]| extract_time_range_strict("ts", TimeUnit::Millisecond, filters);
705 let range = |start: i64, end: i64| {
706 TimestampRange::new(
707 Timestamp::new_millisecond(start),
708 Timestamp::new_millisecond(end),
709 )
710 .unwrap()
711 };
712
713 assert_eq!(extract(&[]), TimeRangeExtraction::Absent);
715 assert_eq!(
716 extract(&[col("host").eq(lit("a"))]),
717 TimeRangeExtraction::Absent
718 );
719
720 assert_eq!(
722 extract(&[
723 col("ts").gt_eq(ts_lit(1000)),
724 col("ts").lt(ts_lit(2000)),
725 col("host").eq(lit("a")),
726 ]),
727 TimeRangeExtraction::Extracted(range(1000, 2000))
728 );
729
730 assert_eq!(
732 extract(&[col("ts").gt_eq(ts_lit(1000))]),
733 TimeRangeExtraction::Extracted(TimestampRange::from_start(Timestamp::new_millisecond(
734 1000
735 )))
736 );
737 assert_eq!(
738 extract(&[col("ts").lt(ts_lit(2000))]),
739 TimeRangeExtraction::Extracted(TimestampRange::until_end(
740 Timestamp::new_millisecond(2000),
741 false
742 ))
743 );
744
745 assert_eq!(
747 extract(&[col("ts").between(ts_lit(1000), ts_lit(2000))]),
748 TimeRangeExtraction::Extracted(range(1000, 2001))
749 );
750
751 assert_eq!(
753 extract(&[col("ts").eq(ts_lit(1500))]),
754 TimeRangeExtraction::Extracted(TimestampRange::single(Timestamp::new_millisecond(
755 1500
756 )))
757 );
758
759 assert_eq!(
761 extract(&[col("ts").gt_eq(ts_lit(1000)).and(col("ts").lt(col("t2")))]),
762 TimeRangeExtraction::Extracted(TimestampRange::from_start(Timestamp::new_millisecond(
763 1000
764 )))
765 );
766
767 let TimeRangeExtraction::Extracted(empty) =
769 extract(&[col("ts").gt_eq(ts_lit(2000)), col("ts").lt(ts_lit(1000))])
770 else {
771 panic!("expected an extraction");
772 };
773 assert!(empty.is_empty());
774
775 assert_eq!(
778 extract(&[col("ts").gt(ts_lit(1000)).or(col("host").eq(lit("a")))]),
779 TimeRangeExtraction::Unsupported
780 );
781 assert_eq!(
782 extract(&[!col("ts").gt(ts_lit(1000))]),
783 TimeRangeExtraction::Unsupported
784 );
785 assert_eq!(
786 extract(&[col("ts").gt_eq(col("t2"))]),
787 TimeRangeExtraction::Unsupported
788 );
789 }
790
791 async fn gen_test_parquet_file(dir: &TempDir, cnt: usize) -> (String, Arc<Schema>) {
792 let path = dir
793 .path()
794 .join("test-prune.parquet")
795 .to_string_lossy()
796 .to_string();
797
798 let name_field = Field::new("name", DataType::Utf8, true);
799 let count_field = Field::new("cnt", DataType::Int32, true);
800 let schema = Arc::new(Schema::new(vec![name_field, count_field]));
801
802 let file = std::fs::OpenOptions::new()
803 .write(true)
804 .create(true)
805 .truncate(true)
806 .open(path.clone())
807 .unwrap();
808
809 let write_props = WriterProperties::builder()
810 .set_max_row_group_row_count(Some(10))
811 .build();
812 let mut writer = ArrowWriter::try_new(file, schema.clone(), Some(write_props)).unwrap();
813
814 for i in (0..cnt).step_by(10) {
815 let name_array = Arc::new(StringArray::from(
816 (i..(i + 10).min(cnt))
817 .map(|i| i.to_string())
818 .collect::<Vec<_>>(),
819 )) as Arc<_>;
820 let count_array = Arc::new(Int32Array::from(
821 (i..(i + 10).min(cnt)).map(|i| i as i32).collect::<Vec<_>>(),
822 )) as Arc<_>;
823 let rb = RecordBatch::try_new(schema.clone(), vec![name_array, count_array]).unwrap();
824 writer.write(&rb).unwrap();
825 }
826 let _ = writer.close().unwrap();
827 (path, schema)
828 }
829
830 async fn assert_prune(array_cnt: usize, filters: Vec<Expr>, expect: Vec<bool>) {
831 let dir = create_temp_dir("prune_parquet");
832 let (path, arrow_schema) = gen_test_parquet_file(&dir, array_cnt).await;
833 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
834 let arrow_predicate = Predicate::new(filters);
835 let builder = ParquetRecordBatchStreamBuilder::new(
836 tokio::fs::OpenOptions::new()
837 .read(true)
838 .open(path)
839 .await
840 .unwrap(),
841 )
842 .await
843 .unwrap();
844 let metadata = builder.metadata().clone();
845 let row_groups = metadata.row_groups();
846
847 let stats = RowGroupPruningStatistics::new(row_groups, &schema);
848 let res = arrow_predicate.prune_with_stats(&stats, &arrow_schema);
849 assert_eq!(expect, res);
850 }
851
852 #[test]
853 fn test_clear_dyn_filters_preserves_static_predicates() {
854 use datafusion_physical_expr::expressions::lit as physical_lit;
855
856 let static_exprs = vec![col("a").eq(lit(1_i32))];
857 let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(vec![], physical_lit(true)));
858 let predicate =
859 Predicate::with_dyn_filters(static_exprs.clone(), vec![dynamic_filter.clone()]);
860
861 predicate.clear_dyn_filters();
862 dynamic_filter.update(physical_lit(false)).unwrap();
864
865 assert_eq!(predicate.exprs(), static_exprs);
866 assert!(predicate.dyn_filters().is_empty());
867 assert!(predicate.dyn_filter_phy_exprs().unwrap().is_empty());
868 }
869
870 #[tokio::test]
871 async fn test_dynamic_pruning_keeps_null_row_group() {
872 use datafusion_physical_expr::expressions::{
873 Column as PhysicalColumn, lit as physical_lit,
874 };
875
876 let dir = create_temp_dir("dynamic_pruning_nulls");
877 let path = dir.path().join("nullable.parquet");
878 let arrow_schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));
879 let file = std::fs::File::create(&path).unwrap();
880 let mut writer = ArrowWriter::try_new(file, arrow_schema.clone(), None).unwrap();
881 for values in [[None, Some(1)], [Some(1), Some(1)], [None, None]] {
882 let batch = RecordBatch::try_new(
883 arrow_schema.clone(),
884 vec![Arc::new(Int32Array::from(values.to_vec()))],
885 )
886 .unwrap();
887 writer.write(&batch).unwrap();
888 writer.flush().unwrap();
889 }
890 writer.close().unwrap();
891 let builder =
892 ParquetRecordBatchStreamBuilder::new(tokio::fs::File::open(path).await.unwrap())
893 .await
894 .unwrap();
895 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
896 let stats = RowGroupPruningStatistics::new(builder.metadata().row_groups(), &schema);
897 let filter = Arc::new(DynamicFilterPhysicalExpr::new(
898 vec![Arc::new(PhysicalColumn::new("a", 0))],
899 physical_lit(true),
900 ));
901 let predicate = Predicate::with_dyn_filters(vec![], vec![filter.clone()]);
902 assert_eq!(
903 predicate.prune_with_stats(&stats, &arrow_schema),
904 vec![true; 3]
905 );
906 filter
907 .update(
908 Predicate::to_physical_expr(
909 &col("a").gt_eq(lit(10_i32)).and(col("a").lt_eq(lit(10_i32))),
910 &arrow_schema,
911 )
912 .unwrap(),
913 )
914 .unwrap();
915 assert_eq!(
917 predicate.prune_with_stats(&stats, &arrow_schema),
918 vec![true, false, true],
919 );
920 }
921
922 #[tokio::test]
923 async fn test_clear_dyn_filters_restores_static_row_group_pruning() {
924 use datafusion_physical_expr::expressions::{
925 Column as PhysicalColumn, lit as physical_lit,
926 };
927
928 let dir = create_temp_dir("dynamic_pruning_reset");
929 let (path, arrow_schema) = gen_test_parquet_file(&dir, 30).await;
930 let schema = Arc::new(datatypes::schema::Schema::try_from(arrow_schema.clone()).unwrap());
931 let builder =
932 ParquetRecordBatchStreamBuilder::new(tokio::fs::File::open(path).await.unwrap())
933 .await
934 .unwrap();
935 let stats = RowGroupPruningStatistics::new(builder.metadata().row_groups(), &schema);
936 let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(
937 vec![Arc::new(PhysicalColumn::new("cnt", 1))],
938 physical_lit(true),
939 ));
940 let predicate = Predicate::with_dyn_filters(
941 vec![col("cnt").gt_eq(lit(10_i32))],
942 vec![dynamic_filter.clone()],
943 );
944
945 dynamic_filter
946 .update(
947 Predicate::to_physical_expr(&col("cnt").gt(lit(100_i32)), &arrow_schema).unwrap(),
948 )
949 .unwrap();
950 assert_eq!(
951 predicate.prune_with_stats(&stats, &arrow_schema),
952 vec![false; 3]
953 );
954
955 predicate.clear_dyn_filters();
958 assert_eq!(
959 predicate.prune_with_stats(&stats, &arrow_schema),
960 vec![false, true, true]
961 );
962 }
963
964 fn gen_predicate(max_val: i32, op: Operator) -> Vec<Expr> {
965 vec![datafusion_expr::Expr::BinaryExpr(BinaryExpr {
966 left: Box::new(datafusion_expr::Expr::Column(Column::from_name("cnt"))),
967 op,
968 right: Box::new(max_val.lit()),
969 })]
970 }
971
972 #[tokio::test]
973 async fn test_prune_empty() {
974 assert_prune(3, vec![], vec![true]).await;
975 }
976
977 #[tokio::test]
978 async fn test_prune_all_match() {
979 let p = gen_predicate(3, Operator::Gt);
980 assert_prune(2, p, vec![false]).await;
981 }
982
983 #[tokio::test]
984 async fn test_prune_gt() {
985 let p = gen_predicate(29, Operator::Gt);
986 assert_prune(
987 100,
988 p,
989 vec![
990 false, false, false, true, true, true, true, true, true, true,
991 ],
992 )
993 .await;
994 }
995
996 #[tokio::test]
997 async fn test_prune_eq_expr() {
998 let p = gen_predicate(30, Operator::Eq);
999 assert_prune(40, p, vec![false, false, false, true]).await;
1000 }
1001
1002 #[tokio::test]
1003 async fn test_prune_neq_expr() {
1004 let p = gen_predicate(30, Operator::NotEq);
1005 assert_prune(40, p, vec![true, true, true, true]).await;
1006 }
1007
1008 #[tokio::test]
1009 async fn test_prune_gteq_expr() {
1010 let p = gen_predicate(29, Operator::GtEq);
1011 assert_prune(40, p, vec![false, false, true, true]).await;
1012 }
1013
1014 #[tokio::test]
1015 async fn test_prune_lt_expr() {
1016 let p = gen_predicate(30, Operator::Lt);
1017 assert_prune(40, p, vec![true, true, true, false]).await;
1018 }
1019
1020 #[tokio::test]
1021 async fn test_prune_lteq_expr() {
1022 let p = gen_predicate(30, Operator::LtEq);
1023 assert_prune(40, p, vec![true, true, true, true]).await;
1024 }
1025
1026 #[tokio::test]
1027 async fn test_prune_between_expr() {
1028 let p = gen_predicate(30, Operator::LtEq);
1029 assert_prune(40, p, vec![true, true, true, true]).await;
1030 }
1031
1032 #[tokio::test]
1033 async fn test_or() {
1034 let e = datafusion_expr::Expr::Column(Column::from_name("cnt"))
1036 .gt(30.lit())
1037 .or(datafusion_expr::Expr::Column(Column::from_name("cnt")).lt(20.lit()));
1038 assert_prune(40, vec![e], vec![true, true, false, true]).await;
1039 }
1040
1041 #[tokio::test]
1042 async fn test_to_physical_expr() {
1043 let predicate = Predicate::new(vec![
1044 col("host").eq(lit("host_a")),
1045 col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(123), None))),
1046 ]);
1047
1048 let schema = Arc::new(arrow::datatypes::Schema::new(vec![Field::new(
1049 "host",
1050 arrow::datatypes::DataType::Utf8,
1051 false,
1052 )]));
1053
1054 let predicates = predicate.to_physical_exprs(&schema).unwrap();
1055 assert!(!predicates.is_empty());
1056
1057 let physical_expr = Predicate::to_physical_expr(&col("host").eq(lit("host_a")), &schema);
1058 assert!(physical_expr.is_ok());
1059 }
1060}