Skip to main content

table/
predicate.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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
44/// Assert the scalar value is not utf8. Returns `None` if it's utf8.
45/// In theory, it should be converted to a timestamp scalar value by `TypeConversionRule`.
46macro_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            // Make the predicate ineffective.
55            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/// Reference-counted pointer to a list of logical exprs and a list of dynamic filter physical exprs.
68#[derive(Debug, Clone, Default)]
69pub struct Predicate {
70    /// logical exprs
71    exprs: Arc<Vec<Expr>>,
72    /// dynamic filter physical exprs, only useful if dynamic filtering is enabled
73    ///
74    /// They are usually from `TopK` or `Join` operators, and can dynamically filter data during query execution by using current runtime information to further reduce data scanning
75    dyn_filters: Arc<ArcSwap<Vec<Arc<DynamicFilterPhysicalExpr>>>>,
76}
77
78impl Predicate {
79    /// Creates a new `Predicate` by converting logical exprs to physical exprs that can be
80    /// evaluated against record batches.
81    /// Returns error when failed to convert exprs.
82    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    /// Adds dynamic filter physical exprs to the existing list.
104    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    /// Removes dynamic filters while preserving the static logical expressions.
113    pub fn clear_dyn_filters(&self) {
114        self.dyn_filters.store(Arc::new(vec![]));
115    }
116
117    /// Returns the logical exprs.
118    pub fn exprs(&self) -> &[Expr] {
119        &self.exprs
120    }
121
122    /// Returns the dynamic filter physical exprs. Notice this return a live dynamic filters which
123    /// can change during query execution.
124    pub fn dyn_filters(&self) -> Arc<Vec<Arc<DynamicFilterPhysicalExpr>>> {
125        self.dyn_filters.load_full()
126    }
127
128    /// Returns the dynamic filter as physical exprs. Notice this return a "snapshot" of
129    /// dynamic filters at the time of calling this method.
130    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                // Pruning must preserve NULL inputs just like decoded dynamic filtering.
136                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    /// Builds a single physical expr according to provided schema.
151    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        // TODO(hl): `execution_props` provides variables required by evaluation.
161        // we may reuse the `execution_props` from `SessionState` once we support
162        // registering variables.
163        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    /// Builds physical exprs according to provided schema.
175    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    /// Evaluates the predicate against the `stats`.
190    /// Returns a vector of boolean values, among which `false` means the row group can be skipped.
191    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                    // since dynamic filter exprs could be complex, it's possible that the pruning predicate builder fails to prove anything from it. In that case, we just log it and skip pruning with this expr.
222                    debug!("Failed to create pruning predicate for expr: {e:?}");
223                }
224            }
225        }
226        res
227    }
228}
229
230// tests for `build_time_range_predicate` locates in src/query/tests/time_range_filter_test.rs
231// since it requires query engine to convert sql to filters.
232/// `build_time_range_predicate` extracts time range from logical exprs to facilitate fast
233/// time range pruning.
234pub 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/// The outcome of strictly extracting the time range of `ts_col_name` from a
249/// scan's filters. Unlike [`build_time_range_predicate`], which quietly widens
250/// to `min_to_max` on anything it cannot parse (fine for pruning), this
251/// distinguishes "the column is not filtered at all" from "it is filtered in a
252/// way that cannot be safely turned into a range" — for callers whose contract
253/// forbids silently falling back to a default window.
254#[derive(Debug, Clone, PartialEq, Eq)]
255pub enum TimeRangeExtraction {
256    /// No filter references the column.
257    Absent,
258    /// Every filter referencing the column was folded into this range, which
259    /// over-approximates the filters' satisfying set.
260    Extracted(TimestampRange),
261    /// At least one filter references the column in a shape that cannot be
262    /// safely extracted (e.g. under `OR`/`NOT`, or compared to a non-literal).
263    Unsupported,
264}
265
266/// Strictly extracts the time range of `ts_col_name` from the (implicitly
267/// AND-ed) scan filters. See [`TimeRangeExtraction`].
268pub 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        // Disjunctive shapes (`OR`, `IN`) collapse disjoint ranges into their
283        // convex hull, which is not exactly representable as one contiguous
284        // range; callers deriving synthetic timestamps from the bounds would
285        // emit points inside the gaps.
286        if contains_disjunction_over_column(expr, ts_col_name) {
287            return TimeRangeExtraction::Unsupported;
288        }
289        // `Some` from the lenient extractor over-approximates the expression's
290        // satisfying set (an `AND` side it cannot parse is dropped, which only
291        // widens), so intersecting extracted conjuncts stays an
292        // over-approximation.
293        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
325/// Extract time range filter from `WHERE`/`IN (...)`/`BETWEEN` clauses.
326/// Return None if no time range can be found in expr.
327pub 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                // [lit] < ts_col
366                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_col < [lit]
373                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                // [lit] <= ts_col
381                ts.convert_to_ceil(ts_col_unit)
382                    .map(TimestampRange::from_start)
383            } else {
384                // ts_col <= [lit]
385                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                // [lit] > ts_col
393                ts.convert_to_ceil(ts_col_unit)
394                    .map(|t| TimestampRange::until_end(t, false))
395            } else {
396                // ts_col > [lit]
397                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                // [lit] >= ts_col
408                ts.convert_to(ts_col_unit)
409                    .map(|t| TimestampRange::until_end(t, true))
410            } else {
411                // ts_col >= [lit]
412                ts.convert_to_ceil(ts_col_unit)
413                    .map(TimestampRange::from_start)
414            }
415        }
416        Operator::And => {
417            // instead of return none when failed to extract time range from left/right, we unwrap the none into
418            // `TimestampRange::min_to_max`.
419            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
484/// Extract time range filter from `IN (...)` expr.
485fn 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                // TODO(hl): maybe we should raise an error here since cannot parse
512                // timestamp value from in list expr
513                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        // ts > 1ms
548        check_build_predicate(
549            col("ts").gt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
550            TimestampRange::from_start(Timestamp::new_millisecond(2)),
551        );
552
553        // 1ms > ts
554        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        // 1001us > ts
560        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        // ts > 1001us
566        check_build_predicate(
567            col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
568            TimestampRange::from_start(Timestamp::new_millisecond(2)),
569        );
570
571        // 1s > ts
572        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        // ts > 1s
578        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        // ts >= 1ms
587        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        // 1ms >= ts
593        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        // 1001us >= ts
599        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        // ts >= 1001us
605        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        // 1s >= ts
611        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        // ts >= 1s
617        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        // ts < 1ms
626        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        // 1ms < ts
632        check_build_predicate(
633            lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt(col("ts")),
634            TimestampRange::from_start(Timestamp::new_millisecond(2)),
635        );
636
637        // 1001us < ts
638        check_build_predicate(
639            lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt(col("ts")),
640            TimestampRange::from_start(Timestamp::new_millisecond(2)),
641        );
642
643        // ts < 1001us
644        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        // 1s < ts
650        check_build_predicate(
651            lit(ScalarValue::TimestampSecond(Some(1), None)).lt(col("ts")),
652            TimestampRange::from_start(Timestamp::new_millisecond(1001)),
653        );
654
655        // ts < 1s
656        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        // ts <= 1ms
665        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        // 1ms <= ts
671        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        // 1001us <= ts
677        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        // ts <= 1001us
683        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        // 1s <= ts
689        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        // ts <= 1s
695        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        // No filter references the column.
717        assert_eq!(extract(&[]), TimeRangeExtraction::Absent);
718        assert_eq!(
719            extract(&[col("host").eq(lit("a"))]),
720            TimeRangeExtraction::Absent
721        );
722
723        // Both bounds across conjuncts; unrelated filters are ignored.
724        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        // Lower bound only / upper bound only.
734        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        // BETWEEN is inclusive on both ends.
749        assert_eq!(
750            extract(&[col("ts").between(ts_lit(1000), ts_lit(2000))]),
751            TimeRangeExtraction::Extracted(range(1000, 2001))
752        );
753
754        // Equality pins a single point.
755        assert_eq!(
756            extract(&[col("ts").eq(ts_lit(1500))]),
757            TimeRangeExtraction::Extracted(TimestampRange::single(Timestamp::new_millisecond(
758                1500
759            )))
760        );
761
762        // An unparsable side under AND only widens; the extraction stays safe.
763        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        // Contradictory bounds collapse to the empty range, not an error.
771        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        // Shapes that could widen the satisfying set beyond what is extractable
779        // must be refused: OR with an unparsable side, NOT, non-literal bounds.
780        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        // An update from the old producer must not put its wrapper back into this execution.
866        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        // NULL inputs must reach decoded filtering; non-NULL misses can still be pruned.
919        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        // Reset after the prior stream is dropped: no dynamic filter remains, while the
959        // static predicate still excludes only the first row group.
960        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        // cnt > 30 or cnt < 20
1038        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}