1use 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#[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
242fn 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 Lossless,
332 TimestampDowncast {
334 source_unit: ArrowTimeUnit,
335 target_unit: ArrowTimeUnit,
336 timezone: Option<Arc<str>>,
337 },
338}
339
340impl NormalizationTarget {
341 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 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 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
437fn 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 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
497fn 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
535fn 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
561fn 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
592fn 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
599fn 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 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}