1use std::sync::Arc;
35
36use arrow::datatypes::{DataType, Schema};
37use datafusion_common::{Result, ScalarValue, tree_node::Transformed};
38use datafusion_expr::Operator;
39use datafusion_expr_common::casts::{
40 is_date_narrowing_cast, is_timestamp_precision_narrowing_cast,
41 try_cast_literal_to_type,
42};
43
44use crate::PhysicalExpr;
45use crate::expressions::{BinaryExpr, CastExpr, Literal, TryCastExpr, lit};
46
47pub(crate) fn unwrap_cast_in_comparison(
49 expr: Arc<dyn PhysicalExpr>,
50 schema: &Schema,
51) -> Result<Transformed<Arc<dyn PhysicalExpr>>> {
52 if let Some(binary) = expr.downcast_ref::<BinaryExpr>()
53 && let Some(unwrapped) = try_unwrap_cast_binary(binary, schema)?
54 {
55 return Ok(Transformed::yes(unwrapped));
56 }
57 Ok(Transformed::no(expr))
58}
59
60fn try_unwrap_cast_binary(
62 binary: &BinaryExpr,
63 schema: &Schema,
64) -> Result<Option<Arc<dyn PhysicalExpr>>> {
65 if let (Some((inner_expr, cast_type)), Some(literal)) = (
67 extract_cast_info(binary.left()),
68 binary.right().downcast_ref::<Literal>(),
69 ) && binary.op().supports_propagation()
70 && let Some(unwrapped) = try_unwrap_cast_comparison(
71 Arc::clone(inner_expr),
72 literal.value(),
73 cast_type,
74 *binary.op(),
75 schema,
76 )?
77 {
78 return Ok(Some(unwrapped));
79 }
80
81 if let (Some(literal), Some((inner_expr, cast_type))) = (
83 binary.left().downcast_ref::<Literal>(),
84 extract_cast_info(binary.right()),
85 ) {
86 if let Some(swapped_op) = binary.op().swap()
88 && binary.op().supports_propagation()
89 && let Some(unwrapped) = try_unwrap_cast_comparison(
90 Arc::clone(inner_expr),
91 literal.value(),
92 cast_type,
93 swapped_op,
94 schema,
95 )?
96 {
97 return Ok(Some(unwrapped));
98 }
99 }
102
103 Ok(None)
104}
105
106fn extract_cast_info(
111 expr: &Arc<dyn PhysicalExpr>,
112) -> Option<(&Arc<dyn PhysicalExpr>, &DataType)> {
113 if let Some(cast) = expr.downcast_ref::<CastExpr>() {
114 Some((cast.expr(), cast.cast_type()))
115 } else if let Some(try_cast) = expr.downcast_ref::<TryCastExpr>() {
116 Some((try_cast.expr(), try_cast.cast_type()))
117 } else {
118 None
119 }
120}
121
122fn try_unwrap_cast_comparison(
124 inner_expr: Arc<dyn PhysicalExpr>,
125 literal_value: &ScalarValue,
126 cast_type: &DataType,
127 op: Operator,
128 schema: &Schema,
129) -> Result<Option<Arc<dyn PhysicalExpr>>> {
130 let inner_type = inner_expr.data_type(schema)?;
132
133 if is_timestamp_precision_narrowing_cast(&inner_type, cast_type)
134 || is_date_narrowing_cast(&inner_type, cast_type)
135 {
136 return Ok(None);
137 }
138
139 if let Some(casted_literal) = try_cast_literal_to_type(literal_value, &inner_type) {
141 let literal_expr = lit(casted_literal);
142 let binary_expr = BinaryExpr::new(inner_expr, op, literal_expr);
143 return Ok(Some(Arc::new(binary_expr)));
144 }
145
146 Ok(None)
147}
148
149#[cfg(test)]
150mod tests {
151 use super::*;
152 use crate::expressions::col;
153 use arrow::datatypes::{Field, TimeUnit};
154 use datafusion_common::tree_node::TreeNode;
155
156 fn is_cast_expr(expr: &Arc<dyn PhysicalExpr>) -> bool {
158 expr.downcast_ref::<CastExpr>().is_some()
159 || expr.downcast_ref::<TryCastExpr>().is_some()
160 }
161
162 fn is_binary_expr_with_cast_and_literal(binary: &BinaryExpr) -> bool {
164 let left_cast_right_literal = is_cast_expr(binary.left())
166 && binary.right().downcast_ref::<Literal>().is_some();
167
168 let left_literal_right_cast = binary.left().downcast_ref::<Literal>().is_some()
170 && is_cast_expr(binary.right());
171
172 left_cast_right_literal || left_literal_right_cast
173 }
174
175 fn test_schema() -> Schema {
176 Schema::new(vec![
177 Field::new("c1", DataType::Int32, false),
178 Field::new("c2", DataType::Int64, false),
179 Field::new("c3", DataType::Utf8, false),
180 ])
181 }
182
183 #[test]
184 fn test_unwrap_cast_in_binary_comparison() {
185 let schema = test_schema();
186
187 let column_expr = col("c1", &schema).unwrap();
189 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
190 let literal_expr = lit(10i64);
191 let binary_expr =
192 Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr));
193
194 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
196
197 assert!(result.transformed);
199
200 let optimized = result.data;
202 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
203
204 assert!(!is_cast_expr(optimized_binary.left()));
206
207 let right_literal = optimized_binary.right().downcast_ref::<Literal>().unwrap();
209 assert_eq!(right_literal.value(), &ScalarValue::Int32(Some(10)));
210 }
211
212 #[test]
213 fn test_unwrap_cast_with_literal_on_left() {
214 let schema = test_schema();
215
216 let column_expr = col("c1", &schema).unwrap();
218 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
219 let literal_expr = lit(10i64);
220 let binary_expr =
221 Arc::new(BinaryExpr::new(literal_expr, Operator::Lt, cast_expr));
222
223 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
225
226 assert!(result.transformed);
228
229 let optimized = result.data;
231 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
232
233 assert_eq!(*optimized_binary.op(), Operator::Gt);
235 }
236
237 #[test]
238 fn test_no_unwrap_date64_to_date32_narrowing() {
239 let schema = Schema::new(vec![Field::new("d64", DataType::Date64, false)]);
240
241 let column_expr = col("d64", &schema).unwrap();
245 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Date32, None));
246 let literal_expr = lit(ScalarValue::Date32(Some(20089)));
247 let binary_expr =
248 Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr));
249
250 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
251 assert!(!result.transformed);
252 }
253
254 #[test]
255 fn test_no_unwrap_when_types_unsupported() {
256 let schema = Schema::new(vec![Field::new("f1", DataType::Float32, false)]);
257
258 let column_expr = col("f1", &schema).unwrap();
260 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Float64, None));
261 let literal_expr = lit(10.5f64);
262 let binary_expr =
263 Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr));
264
265 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
267
268 assert!(!result.transformed);
270 }
271
272 #[test]
273 fn test_is_binary_expr_with_cast_and_literal() {
274 let schema = test_schema();
275
276 let column_expr = col("c1", &schema).unwrap();
277 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
278 let literal_expr = lit(10i64);
279 let binary_expr =
280 Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr));
281 assert!(is_binary_expr_with_cast_and_literal(&binary_expr));
282 }
283
284 #[test]
285 fn test_unwrap_cast_literal_on_left_side() {
286 let schema = Schema::new(vec![Field::new(
289 "decimal_col",
290 DataType::Decimal128(9, 2),
291 true,
292 )]);
293
294 let column_expr = col("decimal_col", &schema).unwrap();
296 let cast_expr = Arc::new(CastExpr::new(
297 column_expr,
298 DataType::Decimal128(22, 2),
299 None,
300 ));
301 let literal_expr = lit(ScalarValue::Decimal128(Some(400), 22, 2));
302 let binary_expr =
303 Arc::new(BinaryExpr::new(literal_expr, Operator::LtEq, cast_expr));
304
305 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
307
308 assert!(result.transformed);
310
311 let optimized = result.data;
313 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
314
315 assert_eq!(*optimized_binary.op(), Operator::GtEq);
317
318 assert!(!is_cast_expr(optimized_binary.left()));
320
321 let right_literal = optimized_binary.right().downcast_ref::<Literal>().unwrap();
323 assert_eq!(
324 right_literal.value().data_type(),
325 DataType::Decimal128(9, 2)
326 );
327 }
328
329 #[test]
330 fn test_unwrap_cast_with_different_comparison_operators() {
331 let schema = Schema::new(vec![Field::new("int_col", DataType::Int32, false)]);
332
333 let operators = vec![
335 (Operator::Lt, Operator::Gt),
336 (Operator::LtEq, Operator::GtEq),
337 (Operator::Gt, Operator::Lt),
338 (Operator::GtEq, Operator::LtEq),
339 (Operator::Eq, Operator::Eq),
340 (Operator::NotEq, Operator::NotEq),
341 ];
342
343 for (original_op, expected_op) in operators {
344 let column_expr = col("int_col", &schema).unwrap();
346 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
347 let literal_expr = lit(100i64);
348 let binary_expr =
349 Arc::new(BinaryExpr::new(literal_expr, original_op, cast_expr));
350
351 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
353
354 assert!(result.transformed);
356
357 let optimized = result.data;
358 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
359
360 assert_eq!(
362 *optimized_binary.op(),
363 expected_op,
364 "Failed for operator {original_op:?} -> {expected_op:?}"
365 );
366
367 assert!(!is_cast_expr(optimized_binary.left()));
369
370 let right_literal =
372 optimized_binary.right().downcast_ref::<Literal>().unwrap();
373 assert_eq!(right_literal.value(), &ScalarValue::Int32(Some(100)));
374 }
375 }
376
377 #[test]
378 fn test_unwrap_cast_with_decimal_types() {
379 let test_cases = vec![
381 (9, 2, 22, 2, 400),
383 (10, 3, 20, 3, 1000),
384 (5, 1, 10, 1, 99),
385 ];
386
387 for (col_p, col_s, cast_p, cast_s, value) in test_cases {
388 let schema = Schema::new(vec![Field::new(
389 "decimal_col",
390 DataType::Decimal128(col_p, col_s),
391 true,
392 )]);
393
394 let column_expr = col("decimal_col", &schema).unwrap();
398 let cast_expr = Arc::new(CastExpr::new(
399 Arc::clone(&column_expr),
400 DataType::Decimal128(cast_p, cast_s),
401 None,
402 ));
403 let literal_expr = lit(ScalarValue::Decimal128(Some(value), cast_p, cast_s));
404 let binary_expr =
405 Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr));
406
407 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
408 assert!(result.transformed);
409
410 let cast_expr = Arc::new(CastExpr::new(
412 column_expr,
413 DataType::Decimal128(cast_p, cast_s),
414 None,
415 ));
416 let literal_expr = lit(ScalarValue::Decimal128(Some(value), cast_p, cast_s));
417 let binary_expr =
418 Arc::new(BinaryExpr::new(literal_expr, Operator::Lt, cast_expr));
419
420 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
421 assert!(result.transformed);
422 }
423 }
424
425 #[test]
426 fn test_unwrap_cast_with_null_literals() {
427 let schema = Schema::new(vec![Field::new("int_col", DataType::Int32, true)]);
429
430 let column_expr = col("int_col", &schema).unwrap();
432 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
433 let null_literal = lit(ScalarValue::Int64(None));
434 let binary_expr =
435 Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, null_literal));
436
437 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
439
440 assert!(result.transformed);
442
443 let optimized = result.data;
445 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
446 let right_literal = optimized_binary.right().downcast_ref::<Literal>().unwrap();
447 assert_eq!(right_literal.value(), &ScalarValue::Int32(None));
448 }
449
450 #[test]
451 fn test_unwrap_cast_with_try_cast() {
452 let schema = Schema::new(vec![Field::new("str_col", DataType::Utf8, true)]);
454
455 let column_expr = col("str_col", &schema).unwrap();
457 let try_cast_expr = Arc::new(TryCastExpr::new(column_expr, DataType::Int64));
458 let literal_expr = lit(100i64);
459 let binary_expr =
460 Arc::new(BinaryExpr::new(try_cast_expr, Operator::Gt, literal_expr));
461
462 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
464
465 assert!(!result.transformed);
467 }
468
469 #[test]
470 fn test_unwrap_cast_preserves_non_comparison_operators() {
471 let schema = Schema::new(vec![Field::new("int_col", DataType::Int32, false)]);
473
474 let column_expr = col("int_col", &schema).unwrap();
476
477 let cast1 = Arc::new(CastExpr::new(
478 Arc::clone(&column_expr),
479 DataType::Int64,
480 None,
481 ));
482 let lit1 = lit(10i64);
483 let compare1 = Arc::new(BinaryExpr::new(cast1, Operator::Gt, lit1));
484
485 let cast2 = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
486 let lit2 = lit(20i64);
487 let compare2 = Arc::new(BinaryExpr::new(cast2, Operator::Lt, lit2));
488
489 let and_expr = Arc::new(BinaryExpr::new(compare1, Operator::And, compare2));
490
491 let result = (and_expr as Arc<dyn PhysicalExpr>)
493 .transform_down(|node| unwrap_cast_in_comparison(node, &schema))
494 .unwrap();
495
496 assert!(result.transformed);
498
499 let optimized = result.data;
501 let and_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
502 assert_eq!(*and_binary.op(), Operator::And);
503
504 let left_binary = and_binary.left().downcast_ref::<BinaryExpr>().unwrap();
506 let right_binary = and_binary.right().downcast_ref::<BinaryExpr>().unwrap();
507
508 assert!(!is_cast_expr(left_binary.left()));
509 assert!(!is_cast_expr(right_binary.left()));
510 }
511
512 #[test]
513 fn test_try_cast_unwrapping() {
514 let schema = test_schema();
515
516 let column_expr = col("c1", &schema).unwrap();
518 let try_cast_expr = Arc::new(TryCastExpr::new(column_expr, DataType::Int64));
519 let literal_expr = lit(100i64);
520 let binary_expr =
521 Arc::new(BinaryExpr::new(try_cast_expr, Operator::LtEq, literal_expr));
522
523 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
525
526 assert!(result.transformed);
528
529 let optimized = result.data;
530 let optimized_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
531
532 assert!(!is_cast_expr(optimized_binary.left()));
534
535 let right_literal = optimized_binary.right().downcast_ref::<Literal>().unwrap();
537 assert_eq!(right_literal.value(), &ScalarValue::Int32(Some(100)));
538 }
539
540 #[test]
541 fn test_non_swappable_operator() {
542 let schema = Schema::new(vec![Field::new("int_col", DataType::Int32, false)]);
544
545 let column_expr = col("int_col", &schema).unwrap();
548 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
549 let literal_expr = lit(10i64);
550 let binary_expr =
551 Arc::new(BinaryExpr::new(literal_expr, Operator::Plus, cast_expr));
552
553 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
555
556 assert!(!result.transformed);
558 }
559
560 #[test]
561 fn test_cast_that_cannot_be_unwrapped_overflow() {
562 let schema = Schema::new(vec![Field::new("small_int", DataType::Int8, false)]);
564
565 let column_expr = col("small_int", &schema).unwrap();
568 let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int64, None));
569 let literal_expr = lit(1000i64); let binary_expr =
571 Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr));
572
573 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
575
576 assert!(!result.transformed);
578 }
579
580 #[test]
581 fn test_not_unwrap_timestamp_precision_narrowing() {
582 let schema = Schema::new(vec![Field::new(
583 "ts",
584 DataType::Timestamp(TimeUnit::Nanosecond, None),
585 false,
586 )]);
587
588 let column_expr = col("ts", &schema).unwrap();
589 let cast_expr = Arc::new(CastExpr::new(
590 column_expr,
591 DataType::Timestamp(TimeUnit::Millisecond, None),
592 None,
593 ));
594 let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1), None));
595 let binary_expr =
596 Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr));
597
598 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
599
600 assert!(!result.transformed);
601 }
602
603 #[test]
604 fn test_unwrap_timestamp_precision_widening() {
605 let schema = Schema::new(vec![Field::new(
606 "ts",
607 DataType::Timestamp(TimeUnit::Millisecond, None),
608 false,
609 )]);
610
611 let column_expr = col("ts", &schema).unwrap();
612 let cast_expr = Arc::new(CastExpr::new(
613 column_expr,
614 DataType::Timestamp(TimeUnit::Nanosecond, None),
615 None,
616 ));
617 let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(1_000_000), None));
618 let binary_expr =
619 Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr));
620
621 let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap();
622
623 assert!(result.transformed);
624 let optimized_binary = result.data.downcast_ref::<BinaryExpr>().unwrap();
625 assert!(!is_cast_expr(optimized_binary.left()));
626 let right_literal = optimized_binary.right().downcast_ref::<Literal>().unwrap();
627 assert_eq!(
628 right_literal.value(),
629 &ScalarValue::TimestampMillisecond(Some(1), None)
630 );
631 }
632
633 #[test]
634 fn test_complex_nested_expression() {
635 let schema = test_schema();
636
637 let c1_expr = col("c1", &schema).unwrap();
640 let c1_cast = Arc::new(CastExpr::new(c1_expr, DataType::Int64, None));
641 let c1_literal = lit(10i64);
642 let c1_binary = Arc::new(BinaryExpr::new(c1_cast, Operator::Gt, c1_literal));
643
644 let c2_expr = col("c2", &schema).unwrap();
645 let c2_cast = Arc::new(CastExpr::new(c2_expr, DataType::Int32, None));
646 let c2_literal = lit(20i32);
647 let c2_binary = Arc::new(BinaryExpr::new(c2_cast, Operator::Eq, c2_literal));
648
649 let and_expr = Arc::new(BinaryExpr::new(c1_binary, Operator::And, c2_binary));
651
652 let result = (and_expr as Arc<dyn PhysicalExpr>)
654 .transform_down(|node| unwrap_cast_in_comparison(node, &schema))
655 .unwrap();
656
657 assert!(result.transformed);
659
660 let optimized = result.data;
662 let and_binary = optimized.downcast_ref::<BinaryExpr>().unwrap();
663
664 let left_binary = and_binary.left().downcast_ref::<BinaryExpr>().unwrap();
666 assert!(!is_cast_expr(left_binary.left()));
667 let left_literal = left_binary.right().downcast_ref::<Literal>().unwrap();
668 assert_eq!(left_literal.value(), &ScalarValue::Int32(Some(10)));
669
670 let right_binary = and_binary.right().downcast_ref::<BinaryExpr>().unwrap();
672 assert!(!is_cast_expr(right_binary.left()));
673 let right_literal = right_binary.right().downcast_ref::<Literal>().unwrap();
674 assert_eq!(right_literal.value(), &ScalarValue::Int64(Some(20)));
675 }
676}