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