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            // TODO(hl): maybe we should raise an error here since cannot parse
509            // timestamp value from in list expr
510            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        // ts > 1ms
545        check_build_predicate(
546            col("ts").gt(lit(ScalarValue::TimestampMillisecond(Some(1), None))),
547            TimestampRange::from_start(Timestamp::new_millisecond(2)),
548        );
549
550        // 1ms > ts
551        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        // 1001us > ts
557        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        // ts > 1001us
563        check_build_predicate(
564            col("ts").gt(lit(ScalarValue::TimestampMicrosecond(Some(1001), None))),
565            TimestampRange::from_start(Timestamp::new_millisecond(2)),
566        );
567
568        // 1s > ts
569        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        // ts > 1s
575        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        // ts >= 1ms
584        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        // 1ms >= ts
590        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        // 1001us >= ts
596        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        // ts >= 1001us
602        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        // 1s >= ts
608        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        // ts >= 1s
614        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        // ts < 1ms
623        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        // 1ms < ts
629        check_build_predicate(
630            lit(ScalarValue::TimestampMillisecond(Some(1), None)).lt(col("ts")),
631            TimestampRange::from_start(Timestamp::new_millisecond(2)),
632        );
633
634        // 1001us < ts
635        check_build_predicate(
636            lit(ScalarValue::TimestampMicrosecond(Some(1001), None)).lt(col("ts")),
637            TimestampRange::from_start(Timestamp::new_millisecond(2)),
638        );
639
640        // ts < 1001us
641        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        // 1s < ts
647        check_build_predicate(
648            lit(ScalarValue::TimestampSecond(Some(1), None)).lt(col("ts")),
649            TimestampRange::from_start(Timestamp::new_millisecond(1001)),
650        );
651
652        // ts < 1s
653        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        // ts <= 1ms
662        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        // 1ms <= ts
668        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        // 1001us <= ts
674        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        // ts <= 1001us
680        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        // 1s <= ts
686        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        // ts <= 1s
692        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        // No filter references the column.
714        assert_eq!(extract(&[]), TimeRangeExtraction::Absent);
715        assert_eq!(
716            extract(&[col("host").eq(lit("a"))]),
717            TimeRangeExtraction::Absent
718        );
719
720        // Both bounds across conjuncts; unrelated filters are ignored.
721        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        // Lower bound only / upper bound only.
731        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        // BETWEEN is inclusive on both ends.
746        assert_eq!(
747            extract(&[col("ts").between(ts_lit(1000), ts_lit(2000))]),
748            TimeRangeExtraction::Extracted(range(1000, 2001))
749        );
750
751        // Equality pins a single point.
752        assert_eq!(
753            extract(&[col("ts").eq(ts_lit(1500))]),
754            TimeRangeExtraction::Extracted(TimestampRange::single(Timestamp::new_millisecond(
755                1500
756            )))
757        );
758
759        // An unparsable side under AND only widens; the extraction stays safe.
760        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        // Contradictory bounds collapse to the empty range, not an error.
768        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        // Shapes that could widen the satisfying set beyond what is extractable
776        // must be refused: OR with an unparsable side, NOT, non-literal bounds.
777        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        // An update from the old producer must not put its wrapper back into this execution.
863        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        // NULL inputs must reach decoded filtering; non-NULL misses can still be pruned.
916        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        // Reset after the prior stream is dropped: no dynamic filter remains, while the
956        // static predicate still excludes only the first row group.
957        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        // cnt > 30 or cnt < 20
1035        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}