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, data_type }) => {
565            Some((CastInputKind::Cast, expr.as_ref(), data_type))
566        }
567        Expr::TryCast(TryCast { expr, data_type }) => {
568            Some((CastInputKind::TryCast, expr.as_ref(), 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                .as_any()
925                .downcast_ref::<FilterExec>()
926                .expect("regex residual must remain a FilterExec");
927            assert!(matches!(
928                filter.schema().field(0).data_type(),
929                DataType::Dictionary(_, value_type) if value_type.as_ref() == &DataType::Utf8
930            ));
931            assert!(!format!("{:?}", filter.predicate()).contains("Cast"));
932
933            let batches = collect(physical_plan, SessionContext::new().task_ctx())
934                .await
935                .unwrap();
936            assert_eq!(
937                expected_rows,
938                batches.iter().map(|batch| batch.num_rows()).sum::<usize>()
939            );
940        }
941    }
942
943    #[test]
944    fn test_dictionary_regex_rewrite_requires_scalar_utf8_pattern() {
945        assert_filter_left_is_cast(
946            vec![
947                Field::new(
948                    "host",
949                    DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
950                    true,
951                ),
952                Field::new("pattern", DataType::Utf8, true),
953            ],
954            Expr::BinaryExpr(BinaryExpr {
955                left: Box::new(cast(col("host"), DataType::Utf8)),
956                op: Operator::RegexMatch,
957                right: Box::new(col("pattern")),
958            }),
959        );
960    }
961
962    #[test]
963    fn test_dictionary_regex_rewrite_excludes_non_regex_and_non_utf8_dictionary() {
964        assert_filter_left_is_cast(
965            vec![Field::new(
966                "host",
967                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
968                true,
969            )],
970            Expr::BinaryExpr(BinaryExpr {
971                left: Box::new(cast(col("host"), DataType::Utf8)),
972                op: Operator::Eq,
973                right: Box::new(lit("api")),
974            }),
975        );
976        assert_filter_left_is_cast(
977            vec![Field::new(
978                "host",
979                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::LargeUtf8)),
980                true,
981            )],
982            Expr::BinaryExpr(BinaryExpr {
983                left: Box::new(cast(col("host"), DataType::Utf8)),
984                op: Operator::RegexMatch,
985                right: Box::new(lit("^api$")),
986            }),
987        );
988    }
989
990    #[test]
991    fn test_dictionary_regex_rewrite_requires_exact_contract() {
992        let dictionary_utf8 = || {
993            vec![Field::new(
994                "host",
995                DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Utf8)),
996                true,
997            )]
998        };
999
1000        for op in [
1001            Operator::RegexMatch,
1002            Operator::RegexIMatch,
1003            Operator::RegexNotMatch,
1004            Operator::RegexNotIMatch,
1005        ] {
1006            let rewritten = rewrite_dictionary_regex(
1007                dictionary_utf8(),
1008                cast(col("host"), DataType::Utf8),
1009                lit("^api$"),
1010                op,
1011            );
1012            assert!(matches!(
1013                rewritten,
1014                Some(Expr::BinaryExpr(BinaryExpr { left, .. })) if matches!(left.as_ref(), Expr::Column(_))
1015            ));
1016        }
1017
1018        for left in [
1019            try_cast(col("host"), DataType::Utf8),
1020            cast(cast(col("host"), DataType::Utf8), DataType::Utf8),
1021        ] {
1022            assert!(
1023                rewrite_dictionary_regex(
1024                    dictionary_utf8(),
1025                    left,
1026                    lit("^api$"),
1027                    Operator::RegexMatch,
1028                )
1029                .is_none()
1030            );
1031        }
1032        assert!(
1033            rewrite_dictionary_regex(
1034                dictionary_utf8(),
1035                cast(col("host"), DataType::Utf8),
1036                lit(ScalarValue::Utf8(None)),
1037                Operator::RegexMatch,
1038            )
1039            .is_none()
1040        );
1041        assert!(
1042            rewrite_dictionary_regex(
1043                vec![Field::new(
1044                    "host",
1045                    DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8)),
1046                    true,
1047                )],
1048                cast(col("host"), DataType::Utf8),
1049                lit("^api$"),
1050                Operator::RegexMatch,
1051            )
1052            .is_none()
1053        );
1054    }
1055
1056    #[test]
1057    fn test_normalize_direct_timestamp_filter() {
1058        assert_timestamp_pushdown(
1059            vec![
1060                Field::new(
1061                    "ts",
1062                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1063                    false,
1064                ),
1065                Field::new("tag", DataType::Utf8, true),
1066            ],
1067            ts_cast_to_ms()
1068                .gt_eq(ts_ms_literal(-299_999))
1069                .and(ts_cast_to_ms().lt_eq(ts_ms_literal(10_000)))
1070                .and(col("tag").eq(lit("api"))),
1071            "Filter: t.ts >= TimestampNanosecond(-299999999999, None) AND t.ts < TimestampNanosecond(10001000000, None) AND t.tag = Utf8(\"api\")\n  TableScan: t",
1072            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999999999, None), t.ts < TimestampNanosecond(10001000000, None), t.tag = Utf8(\"api\")]",
1073            TimestampRange::new_inclusive(
1074                Some(Timestamp::new_nanosecond(-299_999_999_999)),
1075                Some(Timestamp::new_nanosecond(10_000_999_999)),
1076            ),
1077        );
1078    }
1079
1080    #[test]
1081    fn test_normalize_timestamp_between_filter() {
1082        assert_timestamp_pushdown(
1083            vec![Field::new(
1084                "ts",
1085                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1086                false,
1087            )],
1088            ts_cast_to_ms().between(ts_ms_literal(-299_999), ts_ms_literal(10_000)),
1089            "Filter: t.ts >= TimestampNanosecond(-299999999999, None) AND t.ts < TimestampNanosecond(10001000000, None)\n  TableScan: t",
1090            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999999999, None), t.ts < TimestampNanosecond(10001000000, None)]",
1091            TimestampRange::new_inclusive(
1092                Some(Timestamp::new_nanosecond(-299_999_999_999)),
1093                Some(Timestamp::new_nanosecond(10_000_999_999)),
1094            ),
1095        );
1096    }
1097
1098    #[test]
1099    fn test_normalize_strict_timestamp_filter() {
1100        assert_timestamp_pushdown(
1101            vec![Field::new(
1102                "ts",
1103                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1104                false,
1105            )],
1106            ts_cast_to_ms()
1107                .gt(ts_ms_literal(10_000))
1108                .and(ts_cast_to_ms().lt(ts_ms_literal(20_000))),
1109            "Filter: t.ts >= TimestampNanosecond(10001000000, None) AND t.ts < TimestampNanosecond(20000000000, None)\n  TableScan: t",
1110            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(10001000000, None), t.ts < TimestampNanosecond(20000000000, None)]",
1111            TimestampRange::new_inclusive(
1112                Some(Timestamp::new_nanosecond(10_001_000_000)),
1113                Some(Timestamp::new_nanosecond(19_999_999_999)),
1114            ),
1115        );
1116    }
1117
1118    #[test]
1119    fn test_normalize_zero_boundary_timestamp_filter() {
1120        let fields = vec![Field::new(
1121            "ts",
1122            DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1123            false,
1124        )];
1125
1126        assert_timestamp_pushdown(
1127            fields.clone(),
1128            ts_cast_to_ms().gt_eq(ts_ms_literal(0)),
1129            "Filter: t.ts >= TimestampNanosecond(-999999, None)\n  TableScan: t",
1130            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-999999, None)]",
1131            TimestampRange::from_start(Timestamp::new_nanosecond(-999_999)),
1132        );
1133
1134        assert_timestamp_pushdown(
1135            fields.clone(),
1136            ts_cast_to_ms().lt(ts_ms_literal(0)),
1137            "Filter: t.ts < TimestampNanosecond(-999999, None)\n  TableScan: t",
1138            "TableScan: t, full_filters=[t.ts < TimestampNanosecond(-999999, None)]",
1139            TimestampRange::until_end(Timestamp::new_nanosecond(-999_999), false),
1140        );
1141
1142        assert_timestamp_pushdown(
1143            fields,
1144            ts_cast_to_ms().between(ts_ms_literal(0), ts_ms_literal(0)),
1145            "Filter: t.ts >= TimestampNanosecond(-999999, None) AND t.ts < TimestampNanosecond(1000000, None)\n  TableScan: t",
1146            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-999999, None), t.ts < TimestampNanosecond(1000000, None)]",
1147            TimestampRange::new_inclusive(
1148                Some(Timestamp::new_nanosecond(-999_999)),
1149                Some(Timestamp::new_nanosecond(999_999)),
1150            ),
1151        );
1152    }
1153
1154    #[test]
1155    fn test_timestamp_downcast_contract_matches_datafusion_casts() {
1156        let cases = [
1157            (-1_000_001, -1),
1158            (-1_000_000, -1),
1159            (-999_999, 0),
1160            (-1, 0),
1161            (0, 0),
1162            (999_999, 0),
1163            (1_000_000, 1),
1164        ];
1165
1166        for (source, expected) in cases {
1167            let casted = try_cast_literal_to_type(
1168                &ScalarValue::TimestampNanosecond(Some(source), None),
1169                &DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1170            )
1171            .unwrap();
1172            assert_eq!(
1173                ScalarValue::TimestampMillisecond(Some(expected), None),
1174                casted
1175            );
1176        }
1177
1178        assert_eq!(
1179            Some(-1_999_999),
1180            lower_bound_for_ge(-1, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1181        );
1182        assert_eq!(
1183            Some(-999_999),
1184            lower_bound_for_ge(0, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1185        );
1186        assert_eq!(
1187            Some(1_000_000),
1188            lower_bound_for_ge(1, ArrowTimeUnit::Nanosecond, ArrowTimeUnit::Millisecond)
1189        );
1190    }
1191
1192    #[test]
1193    fn test_normalize_plain_timestamp_literals() {
1194        assert_timestamp_pushdown(
1195            vec![Field::new(
1196                "ts",
1197                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1198                false,
1199            )],
1200            col("ts")
1201                .gt_eq(ts_ms_literal(-299_999))
1202                .and(col("ts").lt_eq(ts_ms_literal(10_000))),
1203            "Filter: t.ts >= TimestampNanosecond(-299999000000, None) AND t.ts <= TimestampNanosecond(10000000000, None)\n  TableScan: t",
1204            "TableScan: t, full_filters=[t.ts >= TimestampNanosecond(-299999000000, None), t.ts <= TimestampNanosecond(10000000000, None)]",
1205            TimestampRange::new_inclusive(
1206                Some(Timestamp::new_nanosecond(-299_999_000_000)),
1207                Some(Timestamp::new_nanosecond(10_000_000_000)),
1208            ),
1209        );
1210    }
1211
1212    #[test]
1213    fn test_keep_timestamp_upcast_filter_unchanged() {
1214        assert_filter_plan(
1215            vec![Field::new(
1216                "ts",
1217                DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1218                false,
1219            )],
1220            cast(
1221                col("ts"),
1222                DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1223            )
1224            .gt_eq(lit(ScalarValue::TimestampNanosecond(Some(1), None))),
1225            "Filter: CAST(t.ts AS Timestamp(ns)) >= TimestampNanosecond(1, None)\n  TableScan: t",
1226        );
1227    }
1228
1229    #[test]
1230    fn test_const_normalization_vs_datafusion_cast_preimage_overlap() {
1231        struct Case {
1232            name: &'static str,
1233            fields: Vec<Field>,
1234            predicate: Expr,
1235            expected_greptime: &'static str,
1236            expected_datafusion: &'static str,
1237        }
1238
1239        let cases = [
1240            Case {
1241                name: "integer widening binary",
1242                fields: vec![Field::new("v", DataType::Int16, false)],
1243                predicate: cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
1244                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1245                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1246            },
1247            Case {
1248                name: "swapped integer comparison",
1249                fields: vec![Field::new("v", DataType::Int16, false)],
1250                predicate: lit(42_i64).lt_eq(cast(col("v"), DataType::Int64)),
1251                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1252                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1253            },
1254            Case {
1255                name: "try_cast integer widening binary",
1256                fields: vec![Field::new("v", DataType::Int16, false)],
1257                predicate: try_cast(col("v"), DataType::Int64).gt_eq(lit(42_i64)),
1258                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1259                expected_datafusion: "Filter: t.v >= Int16(42)\n  TableScan: t",
1260            },
1261            Case {
1262                name: "exact in-list",
1263                fields: vec![Field::new("v", DataType::Int16, false)],
1264                predicate: cast(col("v"), DataType::Int64)
1265                    .in_list(vec![lit(1_i64), lit(2_i64)], false),
1266                expected_greptime: "Filter: t.v IN ([Int16(1), Int16(2)])\n  TableScan: t",
1267                expected_datafusion: "Filter: t.v = Int16(1) OR t.v = Int16(2)\n  TableScan: t",
1268            },
1269            Case {
1270                name: "integer between",
1271                fields: vec![Field::new("v", DataType::Int16, false)],
1272                predicate: cast(col("v"), DataType::Int64).between(lit(3_i64), lit(5_i64)),
1273                expected_greptime: "Filter: t.v BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
1274                expected_datafusion: "Filter: t.v >= Int16(3) AND t.v <= Int16(5)\n  TableScan: t",
1275            },
1276            Case {
1277                name: "not between",
1278                fields: vec![Field::new("v", DataType::Int16, false)],
1279                predicate: Expr::Between(Between {
1280                    expr: Box::new(cast(col("v"), DataType::Int64)),
1281                    negated: true,
1282                    low: Box::new(lit(3_i64)),
1283                    high: Box::new(lit(5_i64)),
1284                }),
1285                expected_greptime: "Filter: t.v NOT BETWEEN Int16(3) AND Int16(5)\n  TableScan: t",
1286                expected_datafusion: "Filter: t.v < Int16(3) OR t.v > Int16(5)\n  TableScan: t",
1287            },
1288            Case {
1289                name: "plain literal",
1290                fields: vec![Field::new("v", DataType::Int16, false)],
1291                predicate: col("v").gt_eq(lit(42_i64)),
1292                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1293                expected_datafusion: "Filter: t.v >= Int64(42)\n  TableScan: t",
1294            },
1295            Case {
1296                name: "casted constant",
1297                fields: vec![Field::new("v", DataType::Int16, false)],
1298                predicate: col("v").gt_eq(cast(lit(42_i8), DataType::Int64)),
1299                expected_greptime: "Filter: t.v >= Int16(42)\n  TableScan: t",
1300                expected_datafusion: "Filter: t.v >= Int64(42)\n  TableScan: t",
1301            },
1302            Case {
1303                name: "timestamp downcast equality",
1304                fields: vec![Field::new(
1305                    "ts_ns",
1306                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1307                    false,
1308                )],
1309                predicate: cast(
1310                    col("ts_ns"),
1311                    DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1312                )
1313                .eq(ts_ms_literal(5000)),
1314                expected_greptime: "Filter: CAST(t.ts_ns AS Timestamp(ms)) = TimestampMillisecond(5000, None)\n  TableScan: t",
1315                expected_datafusion: "Filter: t.ts_ns >= TimestampNanosecond(5000000000, None) AND t.ts_ns < TimestampNanosecond(5001000000, None)\n  TableScan: t",
1316            },
1317            Case {
1318                name: "timestamp widening exact",
1319                fields: vec![Field::new(
1320                    "ts_ms",
1321                    DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1322                    false,
1323                )],
1324                predicate: cast(
1325                    col("ts_ms"),
1326                    DataType::Timestamp(ArrowTimeUnit::Nanosecond, None),
1327                )
1328                .eq(lit(ScalarValue::TimestampNanosecond(
1329                    Some(5_000_000_000),
1330                    None,
1331                ))),
1332                expected_greptime: "Filter: CAST(t.ts_ms AS Timestamp(ns)) = TimestampNanosecond(5000000000, None)\n  TableScan: t",
1333                expected_datafusion: "Filter: t.ts_ms = TimestampMillisecond(5000, None)\n  TableScan: t",
1334            },
1335        ];
1336
1337        for case in cases {
1338            let greptime =
1339                greptime_const_normalized_filter(case.fields.clone(), case.predicate.clone());
1340            let datafusion = datafusion_simplified_filter(case.fields, case.predicate);
1341            assert_eq!(case.expected_greptime, greptime, "{} greptime", case.name);
1342            assert_eq!(
1343                case.expected_datafusion, datafusion,
1344                "{} datafusion",
1345                case.name
1346            );
1347        }
1348    }
1349
1350    fn assert_pattern_match_plan(kind: PatternMatchKind, pattern: ScalarValue, expected: &str) {
1351        let predicate = match kind {
1352            PatternMatchKind::Like => Expr::Like(Like::new(
1353                false,
1354                Box::new(cast(col("s"), DataType::LargeUtf8)),
1355                Box::new(lit(pattern)),
1356                None,
1357                false,
1358            )),
1359            PatternMatchKind::SimilarTo => Expr::SimilarTo(Like::new(
1360                false,
1361                Box::new(cast(col("s"), DataType::LargeUtf8)),
1362                Box::new(lit(pattern)),
1363                None,
1364                false,
1365            )),
1366        };
1367
1368        assert_filter_plan(
1369            vec![Field::new("s", DataType::Utf8, false)],
1370            predicate,
1371            expected,
1372        );
1373    }
1374
1375    fn assert_filter_plan(fields: Vec<Field>, predicate: Expr, expected: &str) {
1376        assert_eq!(expected, analyze_filter(fields, predicate).to_string());
1377    }
1378
1379    fn assert_filter_left_is_cast(fields: Vec<Field>, predicate: Expr) {
1380        let analyzed = analyze_filter(fields, predicate);
1381        let LogicalPlan::Filter(filter) = analyzed else {
1382            panic!("expected filter plan");
1383        };
1384        let Expr::BinaryExpr(BinaryExpr { left, .. }) = filter.predicate else {
1385            panic!("expected binary predicate");
1386        };
1387        assert!(matches!(left.as_ref(), Expr::Cast(_)));
1388    }
1389
1390    fn rewrite_dictionary_regex(
1391        fields: Vec<Field>,
1392        left: Expr,
1393        right: Expr,
1394        op: Operator,
1395    ) -> Option<Expr> {
1396        rewrite_dictionary_string_regex(
1397            BinaryExpr {
1398                left: Box::new(left),
1399                op,
1400                right: Box::new(right),
1401            },
1402            &test_schema(fields),
1403        )
1404        .unwrap()
1405    }
1406
1407    fn assert_timestamp_pushdown(
1408        fields: Vec<Field>,
1409        predicate: Expr,
1410        expected_analyzed: &str,
1411        expected_pushed: &str,
1412        expected_range: TimestampRange,
1413    ) {
1414        let analyzed = analyze_filter(fields, predicate);
1415        assert_eq!(expected_analyzed, analyzed.to_string());
1416
1417        let pushed = push_down_filters(analyzed);
1418        assert_eq!(expected_pushed, pushed.to_string());
1419
1420        let range =
1421            build_time_range_predicate("ts", TimeUnit::Nanosecond, &extract_scan_filters(&pushed));
1422        assert_eq!(expected_range, range);
1423    }
1424
1425    fn analyze_filter(fields: Vec<Field>, predicate: Expr) -> LogicalPlan {
1426        analyze_plan(build_filter_plan(test_schema(fields), predicate))
1427    }
1428
1429    fn greptime_const_normalized_filter(fields: Vec<Field>, predicate: Expr) -> String {
1430        analyze_filter(fields, predicate).to_string()
1431    }
1432
1433    fn datafusion_simplified_filter(fields: Vec<Field>, predicate: Expr) -> String {
1434        let plan = build_filter_plan(test_schema(fields), predicate);
1435        Optimizer::with_rules(vec![Arc::new(SimplifyExpressions::new())])
1436            .optimize(plan, &OptimizerContext::new(), |_, _| {})
1437            .unwrap()
1438            .to_string()
1439    }
1440
1441    fn analyze_plan(plan: LogicalPlan) -> LogicalPlan {
1442        ConstNormalizationRule
1443            .analyze(plan, &ConfigOptions::default())
1444            .unwrap()
1445    }
1446
1447    fn build_filter_plan(schema: Arc<DFSchema>, predicate: Expr) -> LogicalPlan {
1448        LogicalPlanBuilder::scan("t", test_source(schema), None)
1449            .unwrap()
1450            .filter(predicate)
1451            .unwrap()
1452            .build()
1453            .unwrap()
1454    }
1455
1456    fn build_scan_plan(schema: Arc<DFSchema>) -> LogicalPlan {
1457        LogicalPlanBuilder::scan("t", test_source(schema), None)
1458            .unwrap()
1459            .build()
1460            .unwrap()
1461    }
1462
1463    fn push_down_filters(plan: LogicalPlan) -> LogicalPlan {
1464        Optimizer::with_rules(vec![Arc::new(PushDownFilter::new())])
1465            .optimize(plan, &OptimizerContext::new(), |_, _| {})
1466            .unwrap()
1467    }
1468
1469    fn ts_cast_to_ms() -> Expr {
1470        cast(
1471            col("ts"),
1472            DataType::Timestamp(ArrowTimeUnit::Millisecond, None),
1473        )
1474    }
1475
1476    fn ts_ms_literal(value: i64) -> Expr {
1477        lit(ScalarValue::TimestampMillisecond(Some(value), None))
1478    }
1479
1480    fn extract_scan_filters(plan: &LogicalPlan) -> Vec<Expr> {
1481        match plan {
1482            LogicalPlan::TableScan(scan) => scan.filters.clone(),
1483            _ => plan
1484                .inputs()
1485                .into_iter()
1486                .flat_map(extract_scan_filters)
1487                .collect(),
1488        }
1489    }
1490
1491    fn test_schema(fields: Vec<Field>) -> Arc<DFSchema> {
1492        arrow_schema::Schema::new(fields).to_dfschema_ref().unwrap()
1493    }
1494
1495    fn test_source(schema: Arc<DFSchema>) -> Arc<dyn TableSource> {
1496        let table = ExactPushdownProvider {
1497            schema: Arc::new(schema.as_ref().as_arrow().clone()),
1498        };
1499        provider_as_source(Arc::new(table))
1500    }
1501
1502    #[derive(Debug)]
1503    struct ExactPushdownProvider {
1504        schema: arrow_schema::SchemaRef,
1505    }
1506
1507    #[async_trait]
1508    impl TableProvider for ExactPushdownProvider {
1509        fn as_any(&self) -> &dyn std::any::Any {
1510            self
1511        }
1512
1513        fn schema(&self) -> arrow_schema::SchemaRef {
1514            self.schema.clone()
1515        }
1516
1517        fn table_type(&self) -> TableType {
1518            TableType::Base
1519        }
1520
1521        async fn scan(
1522            &self,
1523            _state: &dyn Session,
1524            _projection: Option<&Vec<usize>>,
1525            _filters: &[Expr],
1526            _limit: Option<usize>,
1527        ) -> datafusion::error::Result<Arc<dyn ExecutionPlan>> {
1528            unreachable!("scan should not be called in const_normalization tests")
1529        }
1530
1531        fn supports_filters_pushdown(
1532            &self,
1533            filters: &[&Expr],
1534        ) -> datafusion::error::Result<Vec<TableProviderFilterPushDown>> {
1535            Ok(vec![TableProviderFilterPushDown::Exact; filters.len()])
1536        }
1537    }
1538}