Skip to main content

query/optimizer/
const_normalization.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 arrow_schema::{DataType, TimeUnit as ArrowTimeUnit};
18use datafusion::config::ConfigOptions;
19use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion, TreeNodeRewriter};
20use datafusion_common::{DFSchemaRef, Result, ScalarValue};
21use datafusion_expr::expr::{Cast, InList, Like, TryCast};
22use datafusion_expr::{Between, BinaryExpr, Expr, ExprSchemable, LogicalPlan, Operator, lit};
23use datafusion_expr_common::casts::try_cast_literal_to_type;
24use datafusion_optimizer::analyzer::AnalyzerRule;
25
26use crate::plan::ExtractExpr;
27
28/// ConstNormalizationRule rewrites castable constants against their
29/// non-constant comparison operand ahead of filter pushdown.
30#[derive(Debug)]
31pub struct ConstNormalizationRule;
32
33impl AnalyzerRule for ConstNormalizationRule {
34    fn analyze(&self, plan: LogicalPlan, _config: &ConfigOptions) -> Result<LogicalPlan> {
35        plan.transform(|plan| match plan {
36            LogicalPlan::Filter(filter) => {
37                let schema = filter.input.schema().clone();
38                rewrite_plan_exprs(LogicalPlan::Filter(filter), schema)
39            }
40            LogicalPlan::TableScan(scan) => {
41                let schema = scan.projected_schema.clone();
42                rewrite_plan_exprs(LogicalPlan::TableScan(scan), schema)
43            }
44            _ => Ok(Transformed::no(plan)),
45        })
46        .map(|x| x.data)
47    }
48
49    fn name(&self) -> &str {
50        "ConstNormalizationRule"
51    }
52}
53
54fn rewrite_plan_exprs(plan: LogicalPlan, schema: DFSchemaRef) -> Result<Transformed<LogicalPlan>> {
55    let mut rewriter = ConstNormalizationRewriter {
56        schema,
57        transformed: false,
58    };
59    let exprs = plan
60        .expressions_consider_join()
61        .into_iter()
62        .map(|expr| expr.rewrite(&mut rewriter).map(|rewritten| rewritten.data))
63        .collect::<Result<Vec<_>>>()?;
64    if !rewriter.transformed {
65        return Ok(Transformed::no(plan));
66    }
67
68    let inputs = plan.inputs().into_iter().cloned().collect::<Vec<_>>();
69    plan.with_new_exprs(exprs, inputs).map(Transformed::yes)
70}
71
72struct ConstNormalizationRewriter {
73    schema: DFSchemaRef,
74    transformed: bool,
75}
76
77impl TreeNodeRewriter for ConstNormalizationRewriter {
78    type Node = Expr;
79
80    fn f_down(&mut self, expr: Expr) -> Result<Transformed<Expr>> {
81        let recursion = if matches!(
82            expr,
83            Expr::Exists(_) | Expr::InSubquery(_) | Expr::ScalarSubquery(_)
84        ) {
85            TreeNodeRecursion::Jump
86        } else {
87            TreeNodeRecursion::Continue
88        };
89
90        Ok(Transformed::new(expr, false, recursion))
91    }
92
93    fn f_up(&mut self, expr: Expr) -> Result<Transformed<Expr>> {
94        let rewritten = rewrite_expr_node(expr, &self.schema)?;
95        self.transformed |= rewritten.transformed;
96        Ok(rewritten)
97    }
98}
99
100fn rewrite_expr_node(expr: Expr, schema: &DFSchemaRef) -> Result<Transformed<Expr>> {
101    match expr {
102        Expr::BinaryExpr(binary) => match rewrite_binary_expr(binary.clone(), schema)? {
103            Some(expr) => Ok(Transformed::yes(expr)),
104            None => Ok(Transformed::no(Expr::BinaryExpr(binary))),
105        },
106        Expr::Between(between) => match rewrite_between_expr(between.clone(), schema)? {
107            Some(expr) => Ok(Transformed::yes(expr)),
108            None => Ok(Transformed::no(Expr::Between(between))),
109        },
110        Expr::InList(in_list) => match rewrite_in_list_expr(in_list.clone(), schema)? {
111            Some(expr) => Ok(Transformed::yes(expr)),
112            None => Ok(Transformed::no(Expr::InList(in_list))),
113        },
114        Expr::Like(like) => rewrite_like_expr(like, PatternMatchKind::Like, schema),
115        Expr::SimilarTo(like) => rewrite_like_expr(like, PatternMatchKind::SimilarTo, schema),
116        expr => Ok(Transformed::no(expr)),
117    }
118}
119
120fn rewrite_between_expr(between: Between, schema: &DFSchemaRef) -> Result<Option<Expr>> {
121    let Between {
122        expr,
123        negated,
124        low,
125        high,
126    } = between;
127    let expr = *expr;
128    let low_expr = *low;
129    let high_expr = *high;
130    let Some((target, constants)) =
131        extract_rewrite_operands(&expr, &[low_expr.clone(), high_expr.clone()], schema)?
132    else {
133        return Ok(None);
134    };
135
136    if let Some(mut constants) = target.normalize_constants(&constants) {
137        let high = constants
138            .pop()
139            .expect("between normalization expects high constant");
140        let low = constants
141            .pop()
142            .expect("between normalization expects low constant");
143        return Ok(Some(Expr::Between(Between {
144            expr: Box::new(target.expr.clone()),
145            negated,
146            low: Box::new(lit(low)),
147            high: Box::new(lit(high)),
148        })));
149    }
150
151    Ok((!negated)
152        .then(|| target.normalize_timestamp_between(&constants[0], &constants[1]))
153        .flatten())
154}
155
156fn rewrite_in_list_expr(in_list: InList, schema: &DFSchemaRef) -> Result<Option<Expr>> {
157    let InList {
158        expr,
159        list,
160        negated,
161    } = in_list;
162    let expr = *expr;
163    let Some((target, constants)) = extract_rewrite_operands(&expr, &list, schema)? else {
164        return Ok(None);
165    };
166
167    Ok(target.normalize_constants(&constants).map(|constants| {
168        target
169            .expr
170            .clone()
171            .in_list(constants.into_iter().map(lit).collect(), negated)
172    }))
173}
174
175fn rewrite_like_expr(
176    like: Like,
177    kind: PatternMatchKind,
178    schema: &DFSchemaRef,
179) -> Result<Transformed<Expr>> {
180    let original = match kind {
181        PatternMatchKind::Like => Expr::Like(like.clone()),
182        PatternMatchKind::SimilarTo => Expr::SimilarTo(like.clone()),
183    };
184    let Like {
185        negated,
186        expr,
187        pattern,
188        escape_char,
189        case_insensitive,
190    } = like;
191    let expr = *expr;
192    let pattern = *pattern;
193    let Some((target, constants)) =
194        extract_rewrite_operands(&expr, std::slice::from_ref(&pattern), schema)?
195    else {
196        return Ok(Transformed::no(original));
197    };
198    let Some(mut constants) = target.normalize_constants(&constants) else {
199        return Ok(Transformed::no(original));
200    };
201
202    let pattern = lit(constants
203        .pop()
204        .expect("pattern normalization expects one constant"));
205    let like = Like::new(
206        negated,
207        Box::new(target.expr.clone()),
208        Box::new(pattern),
209        escape_char,
210        case_insensitive,
211    );
212    let rewritten = match kind {
213        PatternMatchKind::Like => Expr::Like(like),
214        PatternMatchKind::SimilarTo => Expr::SimilarTo(like),
215    };
216    Ok(Transformed::yes(rewritten))
217}
218
219fn rewrite_binary_expr(binary: BinaryExpr, schema: &DFSchemaRef) -> Result<Option<Expr>> {
220    if let Some(expr) = rewrite_dictionary_string_regex(binary.clone(), schema)? {
221        return Ok(Some(expr));
222    }
223
224    if !binary.op.supports_propagation() {
225        return Ok(None);
226    }
227
228    let BinaryExpr { left, op, right } = binary;
229    let left = *left;
230    let right = *right;
231    if let Some(expr) = rewrite_binary_side(left.clone(), op, right.clone(), schema)? {
232        return Ok(Some(expr));
233    }
234
235    let Some(swapped_op) = op.swap() else {
236        return Ok(None);
237    };
238
239    rewrite_binary_side(right, swapped_op, left, schema)
240}
241
242/// Removes the string coercion/schema-reconciliation cast present before physical regex planning.
243///
244/// Keeping the dictionary input lets DataFusion's physical regex kernel evaluate scalar patterns
245/// against dictionary values instead of materializing the string column first.
246fn rewrite_dictionary_string_regex(
247    binary: BinaryExpr,
248    schema: &DFSchemaRef,
249) -> Result<Option<Expr>> {
250    let BinaryExpr { left, op, right } = binary;
251    if !matches!(
252        &op,
253        Operator::RegexMatch
254            | Operator::RegexIMatch
255            | Operator::RegexNotMatch
256            | Operator::RegexNotIMatch
257    ) || !matches!(right.as_literal(), Some(ScalarValue::Utf8(Some(_))))
258    {
259        return Ok(None);
260    }
261
262    let Some((CastInputKind::Cast, source, DataType::Utf8)) = extract_cast_input(&left) else {
263        return Ok(None);
264    };
265    if !matches!(source, Expr::Column(_))
266        || !matches!(
267            source.get_type(schema)?,
268            DataType::Dictionary(key_type, value_type)
269                if key_type.as_ref() == &DataType::UInt32 && value_type.as_ref() == &DataType::Utf8
270        )
271    {
272        return Ok(None);
273    }
274
275    Ok(Some(Expr::BinaryExpr(BinaryExpr {
276        left: Box::new(source.clone()),
277        op,
278        right,
279    })))
280}
281
282fn rewrite_binary_side(
283    target_expr: Expr,
284    op: Operator,
285    constant_expr: Expr,
286    schema: &DFSchemaRef,
287) -> Result<Option<Expr>> {
288    let Some((target, constants)) =
289        extract_rewrite_operands(&target_expr, std::slice::from_ref(&constant_expr), schema)?
290    else {
291        return Ok(None);
292    };
293
294    if let Some(mut constants) = target.normalize_constants(&constants) {
295        let constant = constants
296            .pop()
297            .expect("binary normalization expects one constant");
298        return Ok(Some(Expr::BinaryExpr(BinaryExpr {
299            left: Box::new(target.expr.clone()),
300            op,
301            right: Box::new(lit(constant)),
302        })));
303    }
304
305    Ok(target.normalize_timestamp_binary(op, &constants[0]))
306}
307
308fn extract_rewrite_operands(
309    target_expr: &Expr,
310    constant_exprs: &[Expr],
311    schema: &DFSchemaRef,
312) -> Result<Option<(NormalizationTarget, Vec<ScalarValue>)>> {
313    let Some(target) = extract_normalization_target(target_expr, schema)? else {
314        return Ok(None);
315    };
316
317    extract_constant_scalars(constant_exprs)
318        .map(|constants| constants.map(|constants| (target, constants)))
319}
320
321#[derive(Clone)]
322struct NormalizationTarget {
323    expr: Expr,
324    data_type: DataType,
325    kind: NormalizationKind,
326}
327
328#[derive(Clone)]
329enum NormalizationKind {
330    /// The cast preserves every source value exactly, so literals can be cast directly.
331    Lossless,
332    /// The cast drops timestamp precision and must widen predicate bounds to preserve semantics.
333    TimestampDowncast {
334        source_unit: ArrowTimeUnit,
335        target_unit: ArrowTimeUnit,
336        timezone: Option<Arc<str>>,
337    },
338}
339
340impl NormalizationTarget {
341    /// Normalizes constants for rewrites that can preserve the original predicate with a direct
342    /// literal cast. Timestamp precision-changing casts are handled by timestamp-specific helpers.
343    fn normalize_constants(&self, constants: &[ScalarValue]) -> Option<Vec<ScalarValue>> {
344        constants
345            .iter()
346            .map(|constant| self.normalize_constant(constant))
347            .collect()
348    }
349
350    fn normalize_constant(&self, constant: &ScalarValue) -> Option<ScalarValue> {
351        match self.kind {
352            NormalizationKind::TimestampDowncast { .. } => None,
353            NormalizationKind::Lossless => try_cast_literal_to_type(constant, &self.data_type),
354        }
355    }
356
357    /// Rewrites predicates over timestamp downcasts into source-side half-open bounds.
358    fn normalize_timestamp_binary(&self, op: Operator, constant: &ScalarValue) -> Option<Expr> {
359        let NormalizationKind::TimestampDowncast {
360            source_unit,
361            target_unit,
362            timezone,
363        } = &self.kind
364        else {
365            return None;
366        };
367
368        let constant = constant
369            .cast_to(&DataType::Timestamp(*target_unit, timezone.clone()))
370            .ok()?;
371        let value = timestamp_scalar_value(&constant)?;
372        let bound = match op {
373            Operator::GtEq => lower_bound_for_ge(value, *source_unit, *target_unit)?,
374            Operator::Gt => lower_bound_for_ge(value.checked_add(1)?, *source_unit, *target_unit)?,
375            Operator::Lt => lower_bound_for_ge(value, *source_unit, *target_unit)?,
376            Operator::LtEq => {
377                lower_bound_for_ge(value.checked_add(1)?, *source_unit, *target_unit)?
378            }
379            _ => return None,
380        };
381
382        let normalized_op = match op {
383            Operator::GtEq | Operator::Gt => Operator::GtEq,
384            Operator::Lt | Operator::LtEq => Operator::Lt,
385            _ => return None,
386        };
387
388        Some(match normalized_op {
389            Operator::GtEq => self.expr.clone().gt_eq(lit(timestamp_scalar(
390                *source_unit,
391                timezone.clone(),
392                bound,
393            ))),
394            Operator::Lt => {
395                self.expr
396                    .clone()
397                    .lt(lit(timestamp_scalar(*source_unit, timezone.clone(), bound)))
398            }
399            _ => unreachable!("timestamp normalization only rewrites to >= or <"),
400        })
401    }
402
403    /// Rewrites `BETWEEN` over timestamp downcasts into an inclusive lower bound and exclusive
404    /// upper bound over the source timestamp unit.
405    fn normalize_timestamp_between(&self, low: &ScalarValue, high: &ScalarValue) -> Option<Expr> {
406        let NormalizationKind::TimestampDowncast {
407            source_unit,
408            target_unit,
409            timezone,
410        } = &self.kind
411        else {
412            return None;
413        };
414
415        let target_type = DataType::Timestamp(*target_unit, timezone.clone());
416        let low = low.cast_to(&target_type).ok()?;
417        let high = high.cast_to(&target_type).ok()?;
418        let low = timestamp_scalar_value(&low)?;
419        let high = timestamp_scalar_value(&high)?;
420
421        let lower = lower_bound_for_ge(low, *source_unit, *target_unit)?;
422        let upper = lower_bound_for_ge(high.checked_add(1)?, *source_unit, *target_unit)?;
423
424        Some(
425            self.expr
426                .clone()
427                .gt_eq(lit(timestamp_scalar(*source_unit, timezone.clone(), lower)))
428                .and(self.expr.clone().lt(lit(timestamp_scalar(
429                    *source_unit,
430                    timezone.clone(),
431                    upper,
432                )))),
433        )
434    }
435}
436
437/// Returns the non-constant side we should normalize against.
438///
439/// Plain expressions normalize literals to their own type. Cast expressions only participate when
440/// the cast is lossless or when timestamp downcasts can be rewritten as wider source-side bounds.
441fn extract_normalization_target(
442    expr: &Expr,
443    schema: &DFSchemaRef,
444) -> Result<Option<NormalizationTarget>> {
445    if extract_constant_scalar(expr)?.is_some() {
446        return Ok(None);
447    }
448
449    let Some((_, source_expr, target_type)) = extract_cast_input(expr) else {
450        return Ok(Some(NormalizationTarget {
451            expr: expr.clone(),
452            data_type: expr.get_type(schema)?,
453            kind: NormalizationKind::Lossless,
454        }));
455    };
456
457    let data_type = source_expr.get_type(schema)?;
458    let Some(kind) = classify_normalization_kind(&data_type, target_type) else {
459        return Ok(None);
460    };
461
462    Ok(Some(NormalizationTarget {
463        expr: source_expr.clone(),
464        data_type,
465        kind,
466    }))
467}
468
469fn classify_normalization_kind(
470    source_type: &DataType,
471    target_type: &DataType,
472) -> Option<NormalizationKind> {
473    // Timestamp casts that change precision need boundary-aware rewrites. A finer target literal
474    // may not map exactly back to the coarser source unit, so the generic lossless path is only
475    // safe for timestamp casts that keep the same unit.
476    if is_lossless_cast(source_type, target_type) {
477        return Some(NormalizationKind::Lossless);
478    }
479
480    match (source_type, target_type) {
481        (
482            DataType::Timestamp(source_unit, source_tz),
483            DataType::Timestamp(target_unit, target_tz),
484        ) if source_tz == target_tz
485            && time_unit_rank(*source_unit) > time_unit_rank(*target_unit) =>
486        {
487            Some(NormalizationKind::TimestampDowncast {
488                source_unit: *source_unit,
489                target_unit: *target_unit,
490                timezone: source_tz.clone(),
491            })
492        }
493        _ => None,
494    }
495}
496
497/// Returns whether every value of `source_type` is representable in `target_type`.
498fn is_lossless_cast(source_type: &DataType, target_type: &DataType) -> bool {
499    match (source_type, target_type) {
500        (DataType::Int8, DataType::Int16 | DataType::Int32 | DataType::Int64)
501        | (DataType::Int16, DataType::Int32 | DataType::Int64)
502        | (DataType::Int32, DataType::Int64)
503        | (DataType::UInt8, DataType::UInt16 | DataType::UInt32 | DataType::UInt64)
504        | (DataType::UInt8, DataType::Int16 | DataType::Int32 | DataType::Int64)
505        | (DataType::UInt16, DataType::UInt32 | DataType::UInt64)
506        | (DataType::UInt16, DataType::Int32 | DataType::Int64)
507        | (DataType::UInt32, DataType::UInt64 | DataType::Int64)
508        | (DataType::Utf8, DataType::Utf8View | DataType::LargeUtf8) => true,
509        (
510            DataType::Timestamp(source_unit, source_tz),
511            DataType::Timestamp(target_unit, target_tz),
512        ) => source_tz == target_tz && source_unit == target_unit,
513        _ => false,
514    }
515}
516
517#[derive(Clone, Copy)]
518enum PatternMatchKind {
519    Like,
520    SimilarTo,
521}
522
523fn extract_constant_scalars(exprs: &[Expr]) -> Result<Option<Vec<ScalarValue>>> {
524    let mut values = Vec::with_capacity(exprs.len());
525    for expr in exprs {
526        let Some(value) = extract_constant_scalar(expr)? else {
527            return Ok(None);
528        };
529        values.push(value);
530    }
531
532    Ok(Some(values))
533}
534
535/// Extracts a literal scalar from an expression, folding constant `CAST` and `TRY_CAST` nodes.
536fn extract_constant_scalar(expr: &Expr) -> Result<Option<ScalarValue>> {
537    if let Some(value) = expr.as_literal() {
538        return Ok(Some(value.clone()));
539    }
540
541    let Some((kind, expr, data_type)) = extract_cast_input(expr) else {
542        return Ok(None);
543    };
544
545    match kind {
546        CastInputKind::Cast => extract_constant_scalar(expr)?
547            .map(|value| value.cast_to(data_type))
548            .transpose(),
549        CastInputKind::TryCast => {
550            Ok(extract_constant_scalar(expr)?.and_then(|value| value.cast_to(data_type).ok()))
551        }
552    }
553}
554
555#[derive(Clone, Copy)]
556enum CastInputKind {
557    Cast,
558    TryCast,
559}
560
561/// Returns the input expression and target type for `CAST` and `TRY_CAST` expressions.
562fn extract_cast_input(expr: &Expr) -> Option<(CastInputKind, &Expr, &DataType)> {
563    match expr {
564        Expr::Cast(Cast { expr, field }) => {
565            Some((CastInputKind::Cast, expr.as_ref(), field.data_type()))
566        }
567        Expr::TryCast(TryCast { expr, field }) => {
568            Some((CastInputKind::TryCast, expr.as_ref(), field.data_type()))
569        }
570        _ => None,
571    }
572}
573
574fn time_unit_rank(unit: ArrowTimeUnit) -> usize {
575    match unit {
576        ArrowTimeUnit::Second => 0,
577        ArrowTimeUnit::Millisecond => 1,
578        ArrowTimeUnit::Microsecond => 2,
579        ArrowTimeUnit::Nanosecond => 3,
580    }
581}
582
583fn time_unit_scale(unit: ArrowTimeUnit) -> i64 {
584    match unit {
585        ArrowTimeUnit::Second => 1,
586        ArrowTimeUnit::Millisecond => 1_000,
587        ArrowTimeUnit::Microsecond => 1_000_000,
588        ArrowTimeUnit::Nanosecond => 1_000_000_000,
589    }
590}
591
592/// Returns the number of source-unit ticks in one target-unit tick for finer-to-coarser casts.
593fn finer_to_coarser_ratio(source_unit: ArrowTimeUnit, target_unit: ArrowTimeUnit) -> Option<i64> {
594    let source_scale = time_unit_scale(source_unit);
595    let target_scale = time_unit_scale(target_unit);
596    (source_scale >= target_scale).then_some(source_scale / target_scale)
597}
598
599/// Returns the smallest source-unit timestamp whose downcast is greater than or equal to
600/// `target_value`.
601///
602/// DataFusion timestamp downcasts truncate toward zero. For non-positive buckets that means the
603/// bucket starts before `target_value * ratio`, so `<= x` can be rewritten as `< lower_bound(x+1)`
604/// without dropping rows near zero or across negative boundaries.
605fn lower_bound_for_ge(
606    target_value: i64,
607    source_unit: ArrowTimeUnit,
608    target_unit: ArrowTimeUnit,
609) -> Option<i64> {
610    let ratio = finer_to_coarser_ratio(source_unit, target_unit)?;
611    let base = target_value.checked_mul(ratio)?;
612    if target_value <= 0 {
613        base.checked_sub(ratio - 1)
614    } else {
615        Some(base)
616    }
617}
618
619fn timestamp_scalar_value(value: &ScalarValue) -> Option<i64> {
620    match value {
621        ScalarValue::TimestampSecond(Some(value), _)
622        | ScalarValue::TimestampMillisecond(Some(value), _)
623        | ScalarValue::TimestampMicrosecond(Some(value), _)
624        | ScalarValue::TimestampNanosecond(Some(value), _) => Some(*value),
625        _ => None,
626    }
627}
628
629fn timestamp_scalar(unit: ArrowTimeUnit, timezone: Option<Arc<str>>, value: i64) -> ScalarValue {
630    match unit {
631        ArrowTimeUnit::Second => ScalarValue::TimestampSecond(Some(value), timezone),
632        ArrowTimeUnit::Millisecond => ScalarValue::TimestampMillisecond(Some(value), timezone),
633        ArrowTimeUnit::Microsecond => ScalarValue::TimestampMicrosecond(Some(value), timezone),
634        ArrowTimeUnit::Nanosecond => ScalarValue::TimestampNanosecond(Some(value), timezone),
635    }
636}
637
638#[cfg(test)]
639mod tests {
640    use std::sync::Arc;
641
642    use arrow::array::{DictionaryArray, StringArray, UInt32Array};
643    use arrow_schema::{DataType, TimeUnit as ArrowTimeUnit};
644    use async_trait::async_trait;
645    use common_time::Timestamp;
646    use common_time::range::TimestampRange;
647    use common_time::timestamp::TimeUnit;
648    use datafusion::catalog::Session;
649    use datafusion::config::ConfigOptions;
650    use datafusion::datasource::{MemTable, TableProvider, provider_as_source};
651    use datafusion::execution::SessionStateBuilder;
652    use datafusion::execution::context::SessionContext;
653    use datafusion::physical_plan::filter::FilterExec;
654    use datafusion::physical_plan::{ExecutionPlan, collect};
655    use datafusion::physical_planner::{DefaultPhysicalPlanner, PhysicalPlanner};
656    use datafusion_common::arrow::datatypes::Field;
657    use datafusion_common::{DFSchema, ScalarValue, ToDFSchema};
658    use datafusion_expr::expr::{Between, BinaryExpr, Like};
659    use datafusion_expr::expr_fn::{cast, col, try_cast};
660    use datafusion_expr::{
661        Expr, LogicalPlan, LogicalPlanBuilder, Operator, TableProviderFilterPushDown, TableScan,
662        TableSource, TableType, lit,
663    };
664    use datafusion_optimizer::analyzer::AnalyzerRule;
665    use datafusion_optimizer::optimizer::{Optimizer, OptimizerContext};
666    use datafusion_optimizer::push_down_filter::PushDownFilter;
667    use datafusion_optimizer::simplify_expressions::SimplifyExpressions;
668    use table::predicate::build_time_range_predicate;
669
670    use super::{
671        ConstNormalizationRule, PatternMatchKind, lower_bound_for_ge,
672        rewrite_dictionary_string_regex, try_cast_literal_to_type,
673    };
674
675    #[test]
676    fn test_normalize_direct_integer_cast_comparison() {
677        assert_filter_plan(
678            vec![Field::new("v", DataType::Int32, false)],
679            cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
680            "Filter: t.v >= Int32(42)\n  TableScan: t",
681        );
682    }
683
684    #[test]
685    fn test_normalize_non_column_operand() {
686        assert_filter_plan(
687            vec![Field::new("v", DataType::Int32, false)],
688            cast(col("v") + lit(1_i32), DataType::Int64).gt_eq(lit(42_i64)),
689            "Filter: t.v + Int32(1) >= Int32(42)\n  TableScan: t",
690        );
691    }
692
693    #[test]
694    fn test_normalize_swapped_binary_comparison() {
695        assert_filter_plan(
696            vec![Field::new("v", DataType::Int16, false)],
697            lit(42_i64).lt_eq(cast(col("v"), DataType::Int64)),
698            "Filter: t.v >= Int16(42)\n  TableScan: t",
699        );
700    }
701
702    #[test]
703    fn test_normalize_try_cast_target() {
704        assert_filter_plan(
705            vec![Field::new("v", DataType::Int16, false)],
706            try_cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
707            "Filter: t.v >= Int16(42)\n  TableScan: t",
708        );
709    }
710
711    #[test]
712    fn test_normalize_casted_constants() {
713        let fields = vec![Field::new("v", DataType::Int16, false)];
714        let cases = [
715            (
716                col("v").gt_eq(cast(lit(42_i8), DataType::Int64)),
717                "Filter: t.v >= Int16(42)\n  TableScan: t",
718            ),
719            (
720                col("v").in_list(
721                    vec![
722                        cast(lit(1_i8), DataType::Int64),
723                        try_cast(lit(2_i8), DataType::Int64),
724                    ],
725                    false,
726                ),
727                "Filter: t.v IN ([Int16(1), Int16(2)])\n  TableScan: t",
728            ),
729        ];
730
731        for (predicate, expected) in cases {
732            assert_filter_plan(fields.clone(), predicate, expected);
733        }
734    }
735
736    #[test]
737    fn test_normalize_plain_integer_literals() {
738        let fields = vec![Field::new("v", DataType::Int16, false)];
739        let cases = [
740            (
741                col("v").gt_eq(lit(42_i64)),
742                "Filter: t.v >= Int16(42)\n  TableScan: t",
743            ),
744            (
745                col("v").in_list(vec![lit(1_i64), lit(2_i64)], false),
746                "Filter: t.v IN ([Int16(1), Int16(2)])\n  TableScan: t",
747            ),
748            (
749                col("v").between(lit(3_i64), lit(5_i64)),
750                "Filter: t.v BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
751            ),
752        ];
753
754        for (predicate, expected) in cases {
755            assert_filter_plan(fields.clone(), predicate, expected);
756        }
757    }
758
759    #[test]
760    fn test_normalize_unsigned_to_signed_literals() {
761        let cases = [
762            (
763                vec![Field::new("v", DataType::UInt8, false)],
764                cast(col("v"), DataType::Int16).lt_eq(lit(255_i16)),
765                "Filter: t.v <= UInt8(255)\n  TableScan: t",
766            ),
767            (
768                vec![Field::new("v", DataType::UInt16, false)],
769                cast(col("v"), DataType::Int32).gt_eq(lit(42_i32)),
770                "Filter: t.v >= UInt16(42)\n  TableScan: t",
771            ),
772            (
773                vec![Field::new("v", DataType::UInt32, false)],
774                cast(col("v"), DataType::Int64).between(lit(3_i64), lit(5_i64)),
775                "Filter: t.v BETWEEN UInt32(3) AND UInt32(5)\n  TableScan: t",
776            ),
777        ];
778
779        for (fields, predicate, expected) in cases {
780            assert_filter_plan(fields, predicate, expected);
781        }
782    }
783
784    #[test]
785    fn test_normalize_in_list_and_between() {
786        let fields = vec![Field::new("v", DataType::Int16, false)];
787        let cases = [
788            (
789                cast(col("v"), DataType::Int64).in_list(vec![lit(1_i64), lit(2_i64)], false),
790                "Filter: t.v IN ([Int16(1), Int16(2)])\n  TableScan: t",
791            ),
792            (
793                cast(col("v"), DataType::Int64).between(lit(3_i64), lit(5_i64)),
794                "Filter: t.v BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
795            ),
796        ];
797
798        for (predicate, expected) in cases {
799            assert_filter_plan(fields.clone(), predicate, expected);
800        }
801    }
802
803    #[test]
804    fn test_keep_non_lossless_literal_unchanged() {
805        assert_filter_plan(
806            vec![Field::new("v", DataType::Int16, false)],
807            col("v").gt_eq(lit(100_000_i64)),
808            "Filter: t.v >= Int64(100000)\n  TableScan: t",
809        );
810    }
811
812    #[test]
813    fn test_normalize_scan_filters() {
814        let scan = build_scan_plan(test_schema(vec![Field::new("v", DataType::Int16, false)]));
815        let LogicalPlan::TableScan(scan) = scan else {
816            panic!("expected table scan");
817        };
818        let plan = LogicalPlan::TableScan(TableScan {
819            filters: vec![cast(col("v"), DataType::Int64).gt_eq(lit(42_i64))],
820            ..scan
821        });
822
823        let analyzed = analyze_plan(plan);
824
825        assert_eq!(
826            vec![col("v").gt_eq(lit(42_i16))],
827            extract_scan_filters(&analyzed)
828        );
829    }
830
831    #[test]
832    fn test_normalize_negated_between() {
833        assert_filter_plan(
834            vec![Field::new("v", DataType::Int16, false)],
835            Expr::Between(Between {
836                expr: Box::new(cast(col("v"), DataType::Int64)),
837                negated: true,
838                low: Box::new(lit(3_i64)),
839                high: Box::new(lit(5_i64)),
840            }),
841            "Filter: t.v NOT BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
842        );
843    }
844
845    #[test]
846    fn test_normalize_like_literal() {
847        assert_pattern_match_plan(
848            PatternMatchKind::Like,
849            ScalarValue::LargeUtf8(Some("api%".to_string())),
850            "Filter: t.s LIKE Utf8(\"api%\")\n  TableScan: t",
851        );
852    }
853
854    #[test]
855    fn test_normalize_similar_to_literal() {
856        assert_pattern_match_plan(
857            PatternMatchKind::SimilarTo,
858            ScalarValue::LargeUtf8(Some("api.*".to_string())),
859            "Filter: t.s SIMILAR TO Utf8(\"api.*\")\n  TableScan: t",
860        );
861    }
862
863    #[tokio::test]
864    async fn test_dictionary_regex_filter_keeps_dictionary_input() {
865        let dictionary_type =
866            DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8));
867        let schema = Arc::new(arrow_schema::Schema::new(vec![Field::new(
868            "host",
869            dictionary_type.clone(),
870            true,
871        )]));
872        let host = DictionaryArray::new(
873            UInt32Array::from(vec![Some(0), Some(1), Some(2), None, Some(3)]),
874            Arc::new(StringArray::from(vec![
875                Some("api"),
876                Some("API"),
877                Some("db"),
878                None,
879            ])),
880        );
881        let batch = datafusion::arrow::record_batch::RecordBatch::try_new(
882            schema.clone(),
883            vec![Arc::new(host)],
884        )
885        .unwrap();
886
887        for (op, expected_rows) in [
888            (Operator::RegexMatch, 1),
889            (Operator::RegexIMatch, 2),
890            (Operator::RegexNotMatch, 2),
891            (Operator::RegexNotIMatch, 1),
892        ] {
893            let table = MemTable::try_new(schema.clone(), vec![vec![batch.clone()]]).unwrap();
894            let predicate = Expr::BinaryExpr(BinaryExpr {
895                // This string coercion/schema-reconciliation cast is present before physical
896                // regex planning and would bypass DataFusion's dictionary-aware scalar regex
897                // kernel.
898                left: Box::new(cast(col("host"), DataType::Utf8)),
899                op,
900                right: Box::new(lit("^api$")),
901            });
902            let plan = LogicalPlanBuilder::scan("t", provider_as_source(Arc::new(table)), None)
903                .unwrap()
904                .filter(predicate)
905                .unwrap()
906                .build()
907                .unwrap();
908            let analyzed = analyze_plan(plan);
909
910            let LogicalPlan::Filter(filter) = &analyzed else {
911                panic!("expected filter plan");
912            };
913            let Expr::BinaryExpr(BinaryExpr { left, .. }) = &filter.predicate else {
914                panic!("expected regex binary predicate");
915            };
916            assert!(matches!(left.as_ref(), Expr::Column(_)));
917
918            let session_state = SessionStateBuilder::new().with_default_features().build();
919            let physical_plan = DefaultPhysicalPlanner::default()
920                .create_physical_plan(&analyzed, &session_state)
921                .await
922                .unwrap();
923            let filter = physical_plan
924                .downcast_ref::<FilterExec>()
925                .expect("regex residual must remain a FilterExec");
926            assert!(matches!(
927                filter.schema().field(0).data_type(),
928                DataType::Dictionary(_, value_type) if value_type.as_ref() == &DataType::Utf8
929            ));
930            assert!(!format!("{:?}", filter.predicate()).contains("Cast"));
931
932            let batches = collect(physical_plan, SessionContext::new().task_ctx())
933                .await
934                .unwrap();
935            assert_eq!(
936                expected_rows,
937                batches.iter().map(|batch| batch.num_rows()).sum::<usize>()
938            );
939        }
940    }
941
942    #[test]
943    fn test_dictionary_regex_rewrite_requires_scalar_utf8_pattern() {
944        assert_filter_left_is_cast(
945            vec![
946                Field::new(
947                    "host",
948                    DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
949                    true,
950                ),
951                Field::new("pattern", DataType::Utf8, true),
952            ],
953            Expr::BinaryExpr(BinaryExpr {
954                left: Box::new(cast(col("host"), DataType::Utf8)),
955                op: Operator::RegexMatch,
956                right: Box::new(col("pattern")),
957            }),
958        );
959    }
960
961    #[test]
962    fn test_dictionary_regex_rewrite_excludes_non_regex_and_non_utf8_dictionary() {
963        assert_filter_left_is_cast(
964            vec![Field::new(
965                "host",
966                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
967                true,
968            )],
969            Expr::BinaryExpr(BinaryExpr {
970                left: Box::new(cast(col("host"), DataType::Utf8)),
971                op: Operator::Eq,
972                right: Box::new(lit("api")),
973            }),
974        );
975        assert_filter_left_is_cast(
976            vec![Field::new(
977                "host",
978                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::LargeUtf8)),
979                true,
980            )],
981            Expr::BinaryExpr(BinaryExpr {
982                left: Box::new(cast(col("host"), DataType::Utf8)),
983                op: Operator::RegexMatch,
984                right: Box::new(lit("^api$")),
985            }),
986        );
987    }
988
989    #[test]
990    fn test_dictionary_regex_rewrite_requires_exact_contract() {
991        let dictionary_utf8 = || {
992            vec![Field::new(
993                "host",
994                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
995                true,
996            )]
997        };
998
999        for op in [
1000            Operator::RegexMatch,
1001            Operator::RegexIMatch,
1002            Operator::RegexNotMatch,
1003            Operator::RegexNotIMatch,
1004        ] {
1005            let rewritten = rewrite_dictionary_regex(
1006                dictionary_utf8(),
1007                cast(col("host"), DataType::Utf8),
1008                lit("^api$"),
1009                op,
1010            );
1011            assert!(matches!(
1012                rewritten,
1013                Some(Expr::BinaryExpr(BinaryExpr { left, .. })) if matches!(left.as_ref(), Expr::Column(_))
1014            ));
1015        }
1016
1017        for left in [
1018            try_cast(col("host"), DataType::Utf8),
1019            cast(cast(col("host"), DataType::Utf8), DataType::Utf8),
1020        ] {
1021            assert!(
1022                rewrite_dictionary_regex(
1023                    dictionary_utf8(),
1024                    left,
1025                    lit("^api$"),
1026                    Operator::RegexMatch,
1027                )
1028                .is_none()
1029            );
1030        }
1031        assert!(
1032            rewrite_dictionary_regex(
1033                dictionary_utf8(),
1034                cast(col("host"), DataType::Utf8),
1035                lit(ScalarValue::Utf8(None)),
1036                Operator::RegexMatch,
1037            )
1038            .is_none()
1039        );
1040        assert!(
1041            rewrite_dictionary_regex(
1042                vec![Field::new(
1043                    "host",
1044                    DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8)),
1045                    true,
1046                )],
1047                cast(col("host"), DataType::Utf8),
1048                lit("^api$"),
1049                Operator::RegexMatch,
1050            )
1051            .is_none()
1052        );
1053    }
1054
1055    #[test]
1056    fn test_normalize_direct_timestamp_filter() {
1057        assert_timestamp_pushdown(
1058            vec![
1059                Field::new(
1060                    "ts",
1061                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1062                    false,
1063                ),
1064                Field::new("tag", DataType::Utf8, true),
1065            ],
1066            ts_cast_to_ms()
1067                .gt_eq(ts_ms_literal(-299_999))
1068                .and(ts_cast_to_ms().lt_eq(ts_ms_literal(10_000)))
1069                .and(col("tag").eq(lit("api"))),
1070            "Filter: t.ts >= TimestampNanosecond(-299999999999, None) AND t.ts < TimestampNanosecond(10001000000, None) AND t.tag = Utf8(\"api\")\n  TableScan: t",
1071            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999999999, None), t.ts < TimestampNanosecond(10001000000, None), t.tag = Utf8(\"api\")]",
1072            TimestampRange::new_inclusive(
1073                Some(Timestamp::new_nanosecond(-299_999_999_999)),
1074                Some(Timestamp::new_nanosecond(10_000_999_999)),
1075            ),
1076        );
1077    }
1078
1079    #[test]
1080    fn test_normalize_timestamp_between_filter() {
1081        assert_timestamp_pushdown(
1082            vec![Field::new(
1083                "ts",
1084                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1085                false,
1086            )],
1087            ts_cast_to_ms().between(ts_ms_literal(-299_999), ts_ms_literal(10_000)),
1088            "Filter: t.ts >= TimestampNanosecond(-299999999999, None) AND t.ts < TimestampNanosecond(10001000000, None)\n  TableScan: t",
1089            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999999999, None), t.ts < TimestampNanosecond(10001000000, None)]",
1090            TimestampRange::new_inclusive(
1091                Some(Timestamp::new_nanosecond(-299_999_999_999)),
1092                Some(Timestamp::new_nanosecond(10_000_999_999)),
1093            ),
1094        );
1095    }
1096
1097    #[test]
1098    fn test_normalize_strict_timestamp_filter() {
1099        assert_timestamp_pushdown(
1100            vec![Field::new(
1101                "ts",
1102                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1103                false,
1104            )],
1105            ts_cast_to_ms()
1106                .gt(ts_ms_literal(10_000))
1107                .and(ts_cast_to_ms().lt(ts_ms_literal(20_000))),
1108            "Filter: t.ts >= TimestampNanosecond(10001000000, None) AND t.ts < TimestampNanosecond(20000000000, None)\n  TableScan: t",
1109            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(10001000000, None), t.ts < TimestampNanosecond(20000000000, None)]",
1110            TimestampRange::new_inclusive(
1111                Some(Timestamp::new_nanosecond(10_001_000_000)),
1112                Some(Timestamp::new_nanosecond(19_999_999_999)),
1113            ),
1114        );
1115    }
1116
1117    #[test]
1118    fn test_normalize_zero_boundary_timestamp_filter() {
1119        let fields = vec![Field::new(
1120            "ts",
1121            DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1122            false,
1123        )];
1124
1125        assert_timestamp_pushdown(
1126            fields.clone(),
1127            ts_cast_to_ms().gt_eq(ts_ms_literal(0)),
1128            "Filter: t.ts >= TimestampNanosecond(-999999, None)\n  TableScan: t",
1129            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-999999, None)]",
1130            TimestampRange::from_start(Timestamp::new_nanosecond(-999_999)),
1131        );
1132
1133        assert_timestamp_pushdown(
1134            fields.clone(),
1135            ts_cast_to_ms().lt(ts_ms_literal(0)),
1136            "Filter: t.ts < TimestampNanosecond(-999999, None)\n  TableScan: t",
1137            "TableScan: t, full_filters=[t.ts < TimestampNanosecond(-999999, None)]",
1138            TimestampRange::until_end(Timestamp::new_nanosecond(-999_999), false),
1139        );
1140
1141        assert_timestamp_pushdown(
1142            fields,
1143            ts_cast_to_ms().between(ts_ms_literal(0), ts_ms_literal(0)),
1144            "Filter: t.ts >= TimestampNanosecond(-999999, None) AND t.ts < TimestampNanosecond(1000000, None)\n  TableScan: t",
1145            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-999999, None), t.ts < TimestampNanosecond(1000000, None)]",
1146            TimestampRange::new_inclusive(
1147                Some(Timestamp::new_nanosecond(-999_999)),
1148                Some(Timestamp::new_nanosecond(999_999)),
1149            ),
1150        );
1151    }
1152
1153    #[test]
1154    fn test_timestamp_downcast_contract_matches_datafusion_casts() {
1155        let cases = [
1156            (-1_000_001, -1),
1157            (-1_000_000, -1),
1158            (-999_999, 0),
1159            (-1, 0),
1160            (0, 0),
1161            (999_999, 0),
1162            (1_000_000, 1),
1163        ];
1164
1165        for (source, expected) in cases {
1166            let casted = try_cast_literal_to_type(
1167                &ScalarValue::TimestampNanosecond(Some(source), None),
1168                &DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1169            )
1170            .unwrap();
1171            assert_eq!(
1172                ScalarValue::TimestampMillisecond(Some(expected), None),
1173                casted
1174            );
1175        }
1176
1177        assert_eq!(
1178            Some(-1_999_999),
1179            lower_bound_for_ge(-1, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1180        );
1181        assert_eq!(
1182            Some(-999_999),
1183            lower_bound_for_ge(0, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1184        );
1185        assert_eq!(
1186            Some(1_000_000),
1187            lower_bound_for_ge(1, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1188        );
1189    }
1190
1191    #[test]
1192    fn test_normalize_plain_timestamp_literals() {
1193        assert_timestamp_pushdown(
1194            vec![Field::new(
1195                "ts",
1196                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1197                false,
1198            )],
1199            col("ts")
1200                .gt_eq(ts_ms_literal(-299_999))
1201                .and(col("ts").lt_eq(ts_ms_literal(10_000))),
1202            "Filter: t.ts >= TimestampNanosecond(-299999000000, None) AND t.ts <= TimestampNanosecond(10000000000, None)\n  TableScan: t",
1203            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999000000, None), t.ts <= TimestampNanosecond(10000000000, None)]",
1204            TimestampRange::new_inclusive(
1205                Some(Timestamp::new_nanosecond(-299_999_000_000)),
1206                Some(Timestamp::new_nanosecond(10_000_000_000)),
1207            ),
1208        );
1209    }
1210
1211    #[test]
1212    fn test_keep_timestamp_upcast_filter_unchanged() {
1213        assert_filter_plan(
1214            vec![Field::new(
1215                "ts",
1216                DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1217                false,
1218            )],
1219            cast(
1220                col("ts"),
1221                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1222            )
1223            .gt_eq(lit(ScalarValue::TimestampNanosecond(Some(1), None))),
1224            "Filter: CAST(t.ts AS Timestamp(ns)) >= TimestampNanosecond(1, None)\n  TableScan: t",
1225        );
1226    }
1227
1228    #[test]
1229    fn test_const_normalization_vs_datafusion_cast_preimage_overlap() {
1230        struct Case {
1231            name: &'static str,
1232            fields: Vec<Field>,
1233            predicate: Expr,
1234            expected_greptime: &'static str,
1235            expected_datafusion: &'static str,
1236        }
1237
1238        let cases = [
1239            Case {
1240                name: "integer widening binary",
1241                fields: vec![Field::new("v", DataType::Int16, false)],
1242                predicate: cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
1243                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1244                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1245            },
1246            Case {
1247                name: "swapped integer comparison",
1248                fields: vec![Field::new("v", DataType::Int16, false)],
1249                predicate: lit(42_i64).lt_eq(cast(col("v"), DataType::Int64)),
1250                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1251                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1252            },
1253            Case {
1254                name: "try_cast integer widening binary",
1255                fields: vec![Field::new("v", DataType::Int16, false)],
1256                predicate: try_cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
1257                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1258                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1259            },
1260            Case {
1261                name: "exact in-list",
1262                fields: vec![Field::new("v", DataType::Int16, false)],
1263                predicate: cast(col("v"), DataType::Int64)
1264                    .in_list(vec![lit(1_i64), lit(2_i64)], false),
1265                expected_greptime: "Filter: t.v IN ([Int16(1), Int16(2)])\n  TableScan: t",
1266                expected_datafusion: "Filter: t.v = Int16(1) OR t.v = Int16(2)\n  TableScan: t",
1267            },
1268            Case {
1269                name: "integer between",
1270                fields: vec![Field::new("v", DataType::Int16, false)],
1271                predicate: cast(col("v"), DataType::Int64).between(lit(3_i64), lit(5_i64)),
1272                expected_greptime: "Filter: t.v BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
1273                expected_datafusion: "Filter: t.v >= Int16(3) AND t.v <= Int16(5)\n  TableScan: t",
1274            },
1275            Case {
1276                name: "not between",
1277                fields: vec![Field::new("v", DataType::Int16, false)],
1278                predicate: Expr::Between(Between {
1279                    expr: Box::new(cast(col("v"), DataType::Int64)),
1280                    negated: true,
1281                    low: Box::new(lit(3_i64)),
1282                    high: Box::new(lit(5_i64)),
1283                }),
1284                expected_greptime: "Filter: t.v NOT BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
1285                expected_datafusion: "Filter: t.v < Int16(3) OR t.v > Int16(5)\n  TableScan: t",
1286            },
1287            Case {
1288                name: "plain literal",
1289                fields: vec![Field::new("v", DataType::Int16, false)],
1290                predicate: col("v").gt_eq(lit(42_i64)),
1291                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1292                expected_datafusion: "Filter: t.v >= Int64(42)\n  TableScan: t",
1293            },
1294            Case {
1295                name: "casted constant",
1296                fields: vec![Field::new("v", DataType::Int16, false)],
1297                predicate: col("v").gt_eq(cast(lit(42_i8), DataType::Int64)),
1298                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1299                expected_datafusion: "Filter: t.v >= Int64(42)\n  TableScan: t",
1300            },
1301            Case {
1302                name: "timestamp downcast equality",
1303                fields: vec![Field::new(
1304                    "ts_ns",
1305                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1306                    false,
1307                )],
1308                predicate: cast(
1309                    col("ts_ns"),
1310                    DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1311                )
1312                .eq(ts_ms_literal(5000)),
1313                expected_greptime: "Filter: CAST(t.ts_ns AS Timestamp(ms)) = TimestampMillisecond(5000, None)\n  TableScan: t",
1314                expected_datafusion: "Filter: t.ts_ns >= TimestampNanosecond(5000000000, None) AND t.ts_ns < TimestampNanosecond(5001000000, None)\n  TableScan: t",
1315            },
1316            Case {
1317                name: "timestamp widening exact",
1318                fields: vec![Field::new(
1319                    "ts_ms",
1320                    DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1321                    false,
1322                )],
1323                predicate: cast(
1324                    col("ts_ms"),
1325                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1326                )
1327                .eq(lit(ScalarValue::TimestampNanosecond(
1328                    Some(5_000_000_000),
1329                    None,
1330                ))),
1331                expected_greptime: "Filter: CAST(t.ts_ms AS Timestamp(ns)) = TimestampNanosecond(5000000000, None)\n  TableScan: t",
1332                expected_datafusion: "Filter: t.ts_ms = TimestampMillisecond(5000, None)\n  TableScan: t",
1333            },
1334            Case {
1335                name: "timestamp widening try_cast exact",
1336                fields: vec![Field::new(
1337                    "ts_ms",
1338                    DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1339                    false,
1340                )],
1341                predicate: try_cast(
1342                    col("ts_ms"),
1343                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1344                )
1345                .eq(lit(ScalarValue::TimestampNanosecond(
1346                    Some(5_000_000_000),
1347                    None,
1348                ))),
1349                expected_greptime: "Filter: TRY_CAST(t.ts_ms AS Timestamp(ns)) = TimestampNanosecond(5000000000, None)\n  TableScan: t",
1350                expected_datafusion: "Filter: t.ts_ms = TimestampMillisecond(5000, None)\n  TableScan: t",
1351            },
1352        ];
1353
1354        for case in cases {
1355            let greptime =
1356                greptime_const_normalized_filter(case.fields.clone(), case.predicate.clone());
1357            let datafusion = datafusion_simplified_filter(case.fields, case.predicate);
1358            assert_eq!(case.expected_greptime, greptime, "{} greptime", case.name);
1359            assert_eq!(
1360                case.expected_datafusion, datafusion,
1361                "{} datafusion",
1362                case.name
1363            );
1364        }
1365    }
1366
1367    fn assert_pattern_match_plan(kind: PatternMatchKind, pattern: ScalarValue, expected: &str) {
1368        let predicate = match kind {
1369            PatternMatchKind::Like => Expr::Like(Like::new(
1370                false,
1371                Box::new(cast(col("s"), DataType::LargeUtf8)),
1372                Box::new(lit(pattern)),
1373                None,
1374                false,
1375            )),
1376            PatternMatchKind::SimilarTo => Expr::SimilarTo(Like::new(
1377                false,
1378                Box::new(cast(col("s"), DataType::LargeUtf8)),
1379                Box::new(lit(pattern)),
1380                None,
1381                false,
1382            )),
1383        };
1384
1385        assert_filter_plan(
1386            vec![Field::new("s", DataType::Utf8, false)],
1387            predicate,
1388            expected,
1389        );
1390    }
1391
1392    fn assert_filter_plan(fields: Vec<Field>, predicate: Expr, expected: &str) {
1393        assert_eq!(expected, analyze_filter(fields, predicate).to_string());
1394    }
1395
1396    fn assert_filter_left_is_cast(fields: Vec<Field>, predicate: Expr) {
1397        let analyzed = analyze_filter(fields, predicate);
1398        let LogicalPlan::Filter(filter) = analyzed else {
1399            panic!("expected filter plan");
1400        };
1401        let Expr::BinaryExpr(BinaryExpr { left, .. }) = filter.predicate else {
1402            panic!("expected binary predicate");
1403        };
1404        assert!(matches!(left.as_ref(), Expr::Cast(_)));
1405    }
1406
1407    fn rewrite_dictionary_regex(
1408        fields: Vec<Field>,
1409        left: Expr,
1410        right: Expr,
1411        op: Operator,
1412    ) -> Option<Expr> {
1413        rewrite_dictionary_string_regex(
1414            BinaryExpr {
1415                left: Box::new(left),
1416                op,
1417                right: Box::new(right),
1418            },
1419            &test_schema(fields),
1420        )
1421        .unwrap()
1422    }
1423
1424    fn assert_timestamp_pushdown(
1425        fields: Vec<Field>,
1426        predicate: Expr,
1427        expected_analyzed: &str,
1428        expected_pushed: &str,
1429        expected_range: TimestampRange,
1430    ) {
1431        let analyzed = analyze_filter(fields, predicate);
1432        assert_eq!(expected_analyzed, analyzed.to_string());
1433
1434        let pushed = push_down_filters(analyzed);
1435        assert_eq!(expected_pushed, pushed.to_string());
1436
1437        let range =
1438            build_time_range_predicate("ts", TimeUnit::Nanosecond, &extract_scan_filters(&pushed));
1439        assert_eq!(expected_range, range);
1440    }
1441
1442    fn analyze_filter(fields: Vec<Field>, predicate: Expr) -> LogicalPlan {
1443        analyze_plan(build_filter_plan(test_schema(fields), predicate))
1444    }
1445
1446    fn greptime_const_normalized_filter(fields: Vec<Field>, predicate: Expr) -> String {
1447        analyze_filter(fields, predicate).to_string()
1448    }
1449
1450    fn datafusion_simplified_filter(fields: Vec<Field>, predicate: Expr) -> String {
1451        let plan = build_filter_plan(test_schema(fields), predicate);
1452        Optimizer::with_rules(vec![Arc::new(SimplifyExpressions::new())])
1453            .optimize(plan, &OptimizerContext::new(), |_, _| {})
1454            .unwrap()
1455            .to_string()
1456    }
1457
1458    fn analyze_plan(plan: LogicalPlan) -> LogicalPlan {
1459        ConstNormalizationRule
1460            .analyze(plan, &ConfigOptions::default())
1461            .unwrap()
1462    }
1463
1464    fn build_filter_plan(schema: Arc<DFSchema>, predicate: Expr) -> LogicalPlan {
1465        LogicalPlanBuilder::scan("t", test_source(schema), None)
1466            .unwrap()
1467            .filter(predicate)
1468            .unwrap()
1469            .build()
1470            .unwrap()
1471    }
1472
1473    fn build_scan_plan(schema: Arc<DFSchema>) -> LogicalPlan {
1474        LogicalPlanBuilder::scan("t", test_source(schema), None)
1475            .unwrap()
1476            .build()
1477            .unwrap()
1478    }
1479
1480    fn push_down_filters(plan: LogicalPlan) -> LogicalPlan {
1481        Optimizer::with_rules(vec![Arc::new(PushDownFilter::new())])
1482            .optimize(plan, &OptimizerContext::new(), |_, _| {})
1483            .unwrap()
1484    }
1485
1486    fn ts_cast_to_ms() -> Expr {
1487        cast(
1488            col("ts"),
1489            DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1490        )
1491    }
1492
1493    fn ts_ms_literal(value: i64) -> Expr {
1494        lit(ScalarValue::TimestampMillisecond(Some(value), None))
1495    }
1496
1497    fn extract_scan_filters(plan: &LogicalPlan) -> Vec<Expr> {
1498        match plan {
1499            LogicalPlan::TableScan(scan) => scan.filters.clone(),
1500            _ => plan
1501                .inputs()
1502                .into_iter()
1503                .flat_map(extract_scan_filters)
1504                .collect(),
1505        }
1506    }
1507
1508    fn test_schema(fields: Vec<Field>) -> Arc<DFSchema> {
1509        arrow_schema::Schema::new(fields).to_dfschema_ref().unwrap()
1510    }
1511
1512    fn test_source(schema: Arc<DFSchema>) -> Arc<dyn TableSource> {
1513        let table = ExactPushdownProvider {
1514            schema: Arc::new(schema.as_ref().as_arrow().clone()),
1515        };
1516        provider_as_source(Arc::new(table))
1517    }
1518
1519    #[derive(Debug)]
1520    struct ExactPushdownProvider {
1521        schema: arrow_schema::SchemaRef,
1522    }
1523
1524    #[async_trait]
1525    impl TableProvider for ExactPushdownProvider {
1526        fn schema(&self) -> arrow_schema::SchemaRef {
1527            self.schema.clone()
1528        }
1529
1530        fn table_type(&self) -> TableType {
1531            TableType::Base
1532        }
1533
1534        async fn scan(
1535            &self,
1536            _state: &dyn Session,
1537            _projection: Option<&Vec<usize>>,
1538            _filters: &[Expr],
1539            _limit: Option<usize>,
1540        ) -> datafusion::error::Result<Arc<dyn ExecutionPlan>> {
1541            unreachable!("scan should not be called in const_normalization tests")
1542        }
1543
1544        fn supports_filters_pushdown(
1545            &self,
1546            filters: &[&Expr],
1547        ) -> datafusion::error::Result<Vec<TableProviderFilterPushDown>> {
1548            Ok(vec![TableProviderFilterPushDown::Exact; filters.len()])
1549        }
1550    }
1551}