From 3d783ebb777c1bd7101ec2d80df0010eacb4a7f1 Mon Sep 17 00:00:00 2001 From: discord9 Date: Thu, 11 Jun 2026 17:47:45 +0800 Subject: [PATCH 01/14] refactor: use cast preimages for cast predicate rewrites --- datafusion/expr-common/src/casts.rs | 363 ++++++++- .../src/simplify_expressions/cast_preimage.rs | 335 ++++++++ .../simplify_expressions/expr_simplifier.rs | 76 +- .../optimizer/src/simplify_expressions/mod.rs | 2 +- .../src/simplify_expressions/unwrap_cast.rs | 712 ------------------ .../optimizer/tests/optimizer_integration.rs | 2 +- .../src/simplifier/unwrap_cast.rs | 132 +++- 7 files changed, 781 insertions(+), 841 deletions(-) create mode 100644 datafusion/optimizer/src/simplify_expressions/cast_preimage.rs delete mode 100644 datafusion/optimizer/src/simplify_expressions/unwrap_cast.rs diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 3518c02772672..6f69aa15429b9 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -22,6 +22,10 @@ //! unwrap_cast module to be shared between logical and physical layers. use std::cmp::Ordering; +use std::sync::Arc; + +use crate::interval_arithmetic::Interval; +use crate::operator::Operator; use arrow::datatypes::{ DataType, MAX_DECIMAL32_FOR_EACH_PRECISION, MAX_DECIMAL64_FOR_EACH_PRECISION, @@ -31,7 +35,18 @@ use arrow::datatypes::{ use arrow::temporal_conversions::{ MICROSECONDS, MILLISECONDS, MILLISECONDS_IN_DAY, NANOSECONDS, }; -use datafusion_common::ScalarValue; +use datafusion_common::{Result, ScalarValue}; + +/// Source-domain preimage of `CAST(source_expr AS target_type) OP literal`. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CastPredicatePreimage { + /// A singleton preimage represented by a literal in the source type. This + /// can keep the original comparison operator. + Exact(ScalarValue), + /// A half-open source-domain interval `[lower, upper)`. The caller must + /// map the comparison operator to range predicates. + Range(Interval), +} /// Convert a literal [`ScalarValue`] to `target_type`, preserving the exact value. /// @@ -76,6 +91,236 @@ pub fn try_cast_literal_to_type( .or_else(|| try_cast_binary(lit_value, target_type)) } +/// Computes a source-domain preimage for `CAST(source AS target_type) OP literal`. +/// +/// This is the shared semantic core for logical and physical cast-predicate +/// rewrites. It returns a singleton [`CastPredicatePreimage::Exact`] for casts +/// where moving the cast to the literal preserves comparison semantics, and a +/// [`CastPredicatePreimage::Range`] for many-to-one casts with known preimages +/// such as timestamp precision narrowing. +pub fn cast_predicate_preimage( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, +) -> Result> { + if let Some(interval) = + timestamp_precision_narrowing_preimage(source_type, target_type, lit_value)? + { + return Ok(Some(CastPredicatePreimage::Range(interval))); + } + + if is_timestamp_precision_narrowing_cast(source_type, target_type) + || is_date_narrowing_cast(source_type, target_type) + { + return Ok(None); + } + + Ok( + exact_preimage_for_cast_predicate(source_type, target_type, op, lit_value) + .map(CastPredicatePreimage::Exact), + ) +} + +/// Computes a singleton source-domain literal for exact cast-predicate rewrites. +/// +/// This intentionally returns `None` for timestamp precision narrowing: those +/// casts are many-to-one and need range preimages instead. +pub fn cast_predicate_exact_literal( + source_type: &DataType, + target_type: &DataType, + lit_value: &ScalarValue, +) -> Option { + if is_timestamp_precision_narrowing_cast(source_type, target_type) + || is_date_narrowing_cast(source_type, target_type) + { + return None; + } + + try_cast_literal_to_type(lit_value, source_type) +} + +/// Returns true when casting a timestamp from `source_type` to `target_type` +/// loses timestamp precision. +pub fn is_timestamp_precision_narrowing_cast( + source_type: &DataType, + target_type: &DataType, +) -> bool { + let (DataType::Timestamp(source_unit, _), DataType::Timestamp(target_unit, _)) = + (source_type, target_type) + else { + return false; + }; + + timestamp_unit_scale(source_unit) > timestamp_unit_scale(target_unit) +} + +fn exact_preimage_for_cast_predicate( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, +) -> Option { + cast_to_string_equality_preimage(source_type, target_type, op, lit_value) + .or_else(|| cast_predicate_exact_literal(source_type, target_type, lit_value)) +} + +/// Computes a singleton preimage for equality predicates over casts whose target +/// value is a string representation of an integer source value. +/// +/// For example, `CAST(int_col AS Utf8) = '123'` can be rewritten to +/// `int_col = 123`, but `CAST(int_col AS Utf8) = '0123'` cannot be rewritten to +/// `int_col = 123` because casting `123` back to a string yields `'123'`, not +/// `'0123'`. +fn cast_to_string_equality_preimage( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, +) -> Option { + if !matches!( + target_type, + DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View + ) { + return None; + } + + match (op, lit_value) { + ( + Operator::Eq | Operator::NotEq, + ScalarValue::Utf8(Some(_)) + | ScalarValue::Utf8View(Some(_)) + | ScalarValue::LargeUtf8(Some(_)), + ) => { + // Only try for integer types (TODO can we do this for other types + // like timestamps)? + use DataType::*; + if matches!( + source_type, + Int8 | Int16 | Int32 | Int64 | UInt8 | UInt16 | UInt32 | UInt64 + ) { + let casted = lit_value.cast_to(source_type).ok()?; + let round_tripped = casted.cast_to(&lit_value.data_type()).ok()?; + if lit_value != &round_tripped { + return None; + } + Some(casted) + } else { + None + } + } + _ => None, + } +} + +fn timestamp_precision_narrowing_preimage( + source_type: &DataType, + target_type: &DataType, + lit_value: &ScalarValue, +) -> Result> { + let ( + DataType::Timestamp(source_unit, source_tz), + DataType::Timestamp(target_unit, _), + ) = (source_type, target_type) + else { + return Ok(None); + }; + + let source_scale = i128::from(timestamp_unit_scale(source_unit)); + let target_scale = i128::from(timestamp_unit_scale(target_unit)); + if source_scale <= target_scale { + return Ok(None); + } + + let Some(target_value) = timestamp_literal_value(lit_value, target_unit) else { + return Ok(None); + }; + + let bucket_width = source_scale / target_scale; + let Some((lower, upper)) = trunc_toward_zero_bucket(target_value, bucket_width) + else { + return Ok(None); + }; + + let Ok(lower) = i64::try_from(lower) else { + return Ok(None); + }; + let Ok(upper) = i64::try_from(upper) else { + return Ok(None); + }; + + Interval::try_new( + timestamp_scalar(source_unit, source_tz.clone(), lower), + timestamp_scalar(source_unit, source_tz.clone(), upper), + ) + .map(Some) +} + +/// Returns the half-open source-domain bucket `[lower, upper)` that truncates +/// toward zero to `value` when divided by `bucket_width`. +/// +/// Timestamp precision narrowing follows integer truncation toward zero rather +/// than mathematical floor. For example, when `bucket_width = 1_000_000`, both +/// `999_999` and `-999_999` truncate to `0`, while `-1_000_000` truncates to +/// `-1`. +/// +/// This makes the inverse bucket depend on the sign of `value`: +/// +/// * `value > 0`: `[value * width, (value + 1) * width)` +/// * `value == 0`: `[1 - width, width)`, spanning small negative and positive +/// values that both truncate to zero +/// * `value < 0`: `[(value - 1) * width + 1, value * width + 1)` +/// +/// The arithmetic uses `checked_*` operations and returns `None` if an +/// intermediate bound cannot be represented as `i128`. +fn trunc_toward_zero_bucket(value: i64, bucket_width: i128) -> Option<(i128, i128)> { + let value = value as i128; + if value > 0 { + let lower = value.checked_mul(bucket_width)?; + let upper = value.checked_add(1)?.checked_mul(bucket_width)?; + Some((lower, upper)) + } else if value == 0 { + Some((1_i128.checked_sub(bucket_width)?, bucket_width)) + } else { + let lower = value + .checked_sub(1)? + .checked_mul(bucket_width)? + .checked_add(1)?; + let upper = value.checked_mul(bucket_width)?.checked_add(1)?; + Some((lower, upper)) + } +} + +fn timestamp_literal_value(lit_value: &ScalarValue, unit: &TimeUnit) -> Option { + match (lit_value, unit) { + (ScalarValue::TimestampSecond(Some(value), _), TimeUnit::Second) + | (ScalarValue::TimestampMillisecond(Some(value), _), TimeUnit::Millisecond) + | (ScalarValue::TimestampMicrosecond(Some(value), _), TimeUnit::Microsecond) + | (ScalarValue::TimestampNanosecond(Some(value), _), TimeUnit::Nanosecond) => { + Some(*value) + } + _ => None, + } +} + +fn timestamp_unit_scale(unit: &TimeUnit) -> i64 { + match unit { + TimeUnit::Second => 1, + TimeUnit::Millisecond => MILLISECONDS, + TimeUnit::Microsecond => MICROSECONDS, + TimeUnit::Nanosecond => NANOSECONDS, + } +} + +fn timestamp_scalar(unit: &TimeUnit, tz: Option>, value: i64) -> ScalarValue { + match unit { + TimeUnit::Second => ScalarValue::TimestampSecond(Some(value), tz), + TimeUnit::Millisecond => ScalarValue::TimestampMillisecond(Some(value), tz), + TimeUnit::Microsecond => ScalarValue::TimestampMicrosecond(Some(value), tz), + TimeUnit::Nanosecond => ScalarValue::TimestampNanosecond(Some(value), tz), + } +} + /// Returns true if unwrap_cast_in_comparison supports this data type pub fn is_supported_type(data_type: &DataType) -> bool { is_supported_numeric_type(data_type) @@ -124,26 +369,6 @@ fn is_lossy_temporal_cast(from_type: &DataType, to_type: &DataType) -> bool { || (is_date_type(to_type) && from_type.is_temporal()) } -/// Returns true when casting a timestamp from `from_type` to `to_type` loses -/// timestamp precision. -/// -/// This is used by comparison cast unwrapping to avoid rewrites such as -/// `CAST(ts_ns AS timestamp(ms)) = lit_ms` -> `ts_ns = lit_ns`. The original -/// predicate can match any nanosecond value in the same millisecond, while the -/// rewritten predicate only matches the exact millisecond boundary. -pub fn is_timestamp_precision_narrowing_cast( - from_type: &DataType, - to_type: &DataType, -) -> bool { - let (DataType::Timestamp(from_unit, _), DataType::Timestamp(to_unit, _)) = - (from_type, to_type) - else { - return false; - }; - - timestamp_unit_scale(from_unit) > timestamp_unit_scale(to_unit) -} - /// Returns true when casting a date column from `from_type` to `to_type` narrows /// `Date64` (milliseconds) to `Date32` (days). /// @@ -156,16 +381,6 @@ pub fn is_timestamp_precision_narrowing_cast( pub fn is_date_narrowing_cast(from_type: &DataType, to_type: &DataType) -> bool { matches!((from_type, to_type), (DataType::Date64, DataType::Date32)) } - -fn timestamp_unit_scale(unit: &TimeUnit) -> i128 { - match unit { - TimeUnit::Second => 1, - TimeUnit::Millisecond => MILLISECONDS as i128, - TimeUnit::Microsecond => MICROSECONDS as i128, - TimeUnit::Nanosecond => NANOSECONDS as i128, - } -} - /// Returns true if unwrap_cast_in_comparison supports this numeric type fn is_supported_numeric_type(data_type: &DataType) -> bool { matches!( @@ -496,18 +711,12 @@ fn try_cast_dictionary( fn cast_between_timestamp(from: &DataType, to: &DataType, value: i128) -> Option { let value = value as i64; let from_scale = match from { - DataType::Timestamp(TimeUnit::Second, _) => 1, - DataType::Timestamp(TimeUnit::Millisecond, _) => MILLISECONDS, - DataType::Timestamp(TimeUnit::Microsecond, _) => MICROSECONDS, - DataType::Timestamp(TimeUnit::Nanosecond, _) => NANOSECONDS, + DataType::Timestamp(unit, _) => timestamp_unit_scale(unit), _ => return Some(value), }; let to_scale = match to { - DataType::Timestamp(TimeUnit::Second, _) => 1, - DataType::Timestamp(TimeUnit::Millisecond, _) => MILLISECONDS, - DataType::Timestamp(TimeUnit::Microsecond, _) => MICROSECONDS, - DataType::Timestamp(TimeUnit::Nanosecond, _) => NANOSECONDS, + DataType::Timestamp(unit, _) => timestamp_unit_scale(unit), _ => return Some(value), }; @@ -1253,6 +1462,82 @@ mod tests { assert_eq!(new_scalar, ScalarValue::TimestampMillisecond(None, None)); } + #[test] + fn test_cast_predicate_preimage_exact() { + assert_eq!( + cast_predicate_preimage( + &DataType::Int32, + &DataType::Int64, + Operator::Gt, + &ScalarValue::Int64(Some(10)), + ) + .unwrap(), + Some(CastPredicatePreimage::Exact(ScalarValue::Int32(Some(10)))) + ); + + assert_eq!( + cast_predicate_preimage( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("123".to_string())), + ) + .unwrap(), + Some(CastPredicatePreimage::Exact(ScalarValue::Int32(Some(123)))) + ); + + assert_eq!( + cast_predicate_preimage( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("0123".to_string())), + ) + .unwrap(), + None + ); + } + + #[test] + fn test_cast_predicate_preimage_timestamp_narrowing_range() { + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(1000), None), + ) + .unwrap(), + Some(CastPredicatePreimage::Range( + Interval::try_new( + ScalarValue::TimestampNanosecond(Some(1_000_000_000), None), + ScalarValue::TimestampNanosecond(Some(1_001_000_000), None), + ) + .unwrap() + )) + ); + + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(-1), None), + ) + .unwrap(), + Some(CastPredicatePreimage::Range( + Interval::try_new( + ScalarValue::TimestampNanosecond(Some(-1_999_999), None), + ScalarValue::TimestampNanosecond(Some(-999_999), None), + ) + .unwrap() + )) + ); + } + #[test] fn test_try_cast_to_string_type() { let scalars = vec![ diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs new file mode 100644 index 0000000000000..aaf3efb5e612b --- /dev/null +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -0,0 +1,335 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Preimage rewrites for cast comparisons. +//! +//! This module computes source-domain predicates for expressions such as +//! `CAST(expr AS target_type) OP literal`. For casts that are many-to-one, a +//! same-operator unwrap is not equivalent; the correct rewrite is a preimage +//! range over the input expression. + +use arrow::datatypes::DataType; +use datafusion_common::{Result, internal_err, tree_node::Transformed}; +use datafusion_expr::expr::InList; +use datafusion_expr::{ + BinaryExpr, Cast, Expr, Operator, TryCast, lit, simplify::SimplifyContext, +}; +use datafusion_expr_common::casts::{ + CastPredicatePreimage, cast_predicate_exact_literal, cast_predicate_preimage, +}; + +use super::udf_preimage::rewrite_with_preimage; + +pub(super) fn rewrite_cast_predicate_for_binary( + info: &SimplifyContext, + cast_expr: Expr, + literal: Expr, + op: Operator, +) -> Result> { + let Some((expr, target_type)) = cast_input_and_type(cast_expr) else { + return internal_err!("Expect cast expr"); + }; + let Expr::Literal(lit_value, _) = literal else { + return internal_err!("Expect literal expr"); + }; + + let source_type = info.get_data_type(&expr)?; + match cast_predicate_preimage(&source_type, &target_type, op, &lit_value)? { + Some(CastPredicatePreimage::Range(interval)) => { + rewrite_with_preimage(interval, op, *expr) + } + Some(CastPredicatePreimage::Exact(value)) => { + Ok(Transformed::yes(Expr::BinaryExpr(BinaryExpr { + left: expr, + op, + right: Box::new(lit(value)), + }))) + } + None => internal_err!( + "Can't compute cast predicate preimage for source type {} target type {} literal {:?}", + source_type, + target_type, + lit_value + ), + } +} + +pub(super) fn supports_cast_predicate_for_binary( + info: &SimplifyContext, + expr: &Expr, + op: Operator, + literal: &Expr, +) -> bool { + if !matches!( + op, + Operator::Eq + | Operator::NotEq + | Operator::Lt + | Operator::LtEq + | Operator::Gt + | Operator::GtEq + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom + ) { + return false; + } + + let Some((inner_expr, target_type)) = cast_input_and_type_ref(expr) else { + return false; + }; + let Expr::Literal(lit_value, _) = literal else { + return false; + }; + let Ok(source_type) = info.get_data_type(inner_expr) else { + return false; + }; + + cast_predicate_preimage(&source_type, target_type, op, lit_value) + .ok() + .flatten() + .is_some() +} + +pub(super) fn supports_cast_predicate_for_inlist( + info: &SimplifyContext, + expr: &Expr, + list: &[Expr], +) -> bool { + let Some((inner_expr, target_type)) = cast_input_and_type_ref(expr) else { + return false; + }; + let Ok(source_type) = info.get_data_type(inner_expr) else { + return false; + }; + + list.iter().all(|right| match right { + Expr::Literal(lit_val, _) => { + cast_predicate_exact_literal(&source_type, target_type, lit_val).is_some() + } + _ => false, + }) +} + +pub(super) fn rewrite_cast_predicate_for_inlist( + info: &SimplifyContext, + expr: Expr, + list: Vec, + negated: bool, +) -> Result> { + let Some((inner_expr, _target_type)) = cast_input_and_type(expr) else { + return internal_err!("Expect cast expr"); + }; + let source_type = info.get_data_type(&inner_expr)?; + + let list = list + .into_iter() + .map(|right| match right { + Expr::Literal(lit_value, _) => { + let Some(value) = + cast_predicate_exact_literal(&source_type, &_target_type, &lit_value) + else { + return internal_err!( + "Can't cast the list expr {:?} to type {}", + lit_value, + &source_type + ); + }; + Ok(lit(value)) + } + other_expr => internal_err!( + "Only support literal expr to optimize, but the expr is {:?}", + &other_expr + ), + }) + .collect::>>()?; + + Ok(Transformed::yes(Expr::InList(InList { + expr: inner_expr, + list, + negated, + }))) +} + +fn cast_input_and_type(cast_expr: Expr) -> Option<(Box, DataType)> { + match cast_expr { + Expr::TryCast(TryCast { expr, field, .. }) + | Expr::Cast(Cast { expr, field, .. }) => Some((expr, field.data_type().clone())), + _ => None, + } +} + +fn cast_input_and_type_ref(cast_expr: &Expr) -> Option<(&Expr, &DataType)> { + match cast_expr { + Expr::TryCast(TryCast { expr, field, .. }) + | Expr::Cast(Cast { expr, field, .. }) => { + Some((expr.as_ref(), field.data_type())) + } + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + use std::sync::Arc; + + use crate::simplify_expressions::ExprSimplifier; + use arrow::datatypes::{Field, TimeUnit}; + use datafusion_common::{DFSchema, DFSchemaRef, ScalarValue}; + use datafusion_expr::simplify::SimplifyContext; + use datafusion_expr::{cast, col, in_list}; + + #[test] + fn test_cast_predicate_exact_literal_unwrap() { + let schema = expr_test_schema(); + + let expr = cast(col("c1"), DataType::Int64).gt(lit(10_i64)); + let expected = col("c1").gt(lit(10_i32)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = lit(10_i64).lt(cast(col("c1"), DataType::Int64)); + let expected = col("c1").gt(lit(10_i32)); + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_cast_predicate_string_integer_round_trip() { + let schema = expr_test_schema(); + + let expr = cast(col("c1"), DataType::Utf8).eq(lit("123")); + let expected = col("c1").eq(lit(123_i32)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = cast(col("c1"), DataType::Utf8).eq(lit("0123")); + assert_eq!(optimize_test(expr.clone(), &schema), expr); + } + + #[test] + fn test_cast_predicate_inlist_exact_unwrap() { + let schema = expr_test_schema(); + + let expr = in_list( + cast(col("c1"), DataType::Int64), + vec![lit(0_i64), lit(1_i64), lit(2_i64), lit(3_i64), lit(4_i64)], + false, + ); + let expected = in_list( + col("c1"), + vec![lit(0_i32), lit(1_i32), lit(2_i32), lit(3_i32), lit(4_i32)], + false, + ); + + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_no_unwrap_date64_to_date32_narrowing() { + let schema = Arc::new( + DFSchema::from_unqualified_fields( + vec![Field::new("d64", DataType::Date64, false)].into(), + HashMap::new(), + ) + .unwrap(), + ); + let expr = + cast(col("d64"), DataType::Date32).eq(lit(ScalarValue::Date32(Some(20089)))); + + assert_eq!(optimize_test(expr.clone(), &schema), expr); + } + + #[test] + fn test_cast_preimage_timestamp_precision_narrowing_eq() { + let schema = expr_test_schema(); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).eq(lit_timestamp_millis(1000)); + let expected = col("ts_nano") + .gt_eq(lit_timestamp_nano(1_000_000_000)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_001_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).eq(lit_timestamp_millis(0)); + let expected = col("ts_nano") + .gt_eq(lit_timestamp_nano(-999_999)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).eq(lit_timestamp_millis(-1)); + let expected = col("ts_nano") + .gt_eq(lit_timestamp_nano(-1_999_999)) + .and(col("ts_nano").lt(lit_timestamp_nano(-999_999))); + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_cast_preimage_timestamp_precision_narrowing_inequality() { + let schema = expr_test_schema(); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).gt(lit_timestamp_millis(1000)); + let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).lt_eq(lit_timestamp_millis(-1)); + let expected = col("ts_nano").lt(lit_timestamp_nano(-999_999)); + assert_eq!(optimize_test(expr, &schema), expected); + } + + fn optimize_test(expr: Expr, schema: &DFSchemaRef) -> Expr { + let simplifier = ExprSimplifier::new( + SimplifyContext::builder() + .with_schema(Arc::clone(schema)) + .build(), + ); + + simplifier.simplify(expr).unwrap() + } + + fn expr_test_schema() -> DFSchemaRef { + Arc::new( + DFSchema::from_unqualified_fields( + vec![ + Field::new("c1", DataType::Int32, false), + Field::new("ts_nano", timestamp_nano_type(), false), + ] + .into(), + HashMap::new(), + ) + .unwrap(), + ) + } + + fn lit_timestamp_nano(ts: i64) -> Expr { + lit(ScalarValue::TimestampNanosecond(Some(ts), None)) + } + + fn lit_timestamp_millis(ts: i64) -> Expr { + lit(ScalarValue::TimestampMillisecond(Some(ts), None)) + } + + fn timestamp_nano_type() -> DataType { + DataType::Timestamp(TimeUnit::Nanosecond, None) + } + + fn timestamp_millis_type() -> DataType { + DataType::Timestamp(TimeUnit::Millisecond, None) + } +} diff --git a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs index 5436bd092163e..cfb776e7f2232 100644 --- a/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs +++ b/datafusion/optimizer/src/simplify_expressions/expr_simplifier.rs @@ -59,18 +59,16 @@ use datafusion_physical_expr::{create_physical_expr, execution_props::ExecutionP use super::inlist_simplifier::ShortenInListSimplifier; use super::utils::*; use crate::simplify_expressions::SimplifyContext; -use crate::simplify_expressions::regex::simplify_regex_expr; -use crate::simplify_expressions::unwrap_cast::{ - is_cast_expr_and_support_unwrap_cast_in_comparison_for_binary, - is_cast_expr_and_support_unwrap_cast_in_comparison_for_inlist, - unwrap_cast_in_comparison_for_binary, +use crate::simplify_expressions::cast_preimage::{ + rewrite_cast_predicate_for_binary, rewrite_cast_predicate_for_inlist, + supports_cast_predicate_for_binary, supports_cast_predicate_for_inlist, }; +use crate::simplify_expressions::regex::simplify_regex_expr; use crate::{ analyzer::type_coercion::TypeCoercionRewriter, simplify_expressions::udf_preimage::rewrite_with_preimage, }; use datafusion_expr::expr_rewriter::rewrite_with_guarantees_map; -use datafusion_expr_common::casts::try_cast_literal_to_type; use indexmap::IndexSet; use regex::{Error as RegexError, Regex, RegexBuilder}; @@ -2014,28 +2012,27 @@ impl TreeNodeRewriter for Simplifier<'_> { } // ======================================= - // unwrap_cast_in_comparison + // cast_predicate_in_comparison // ======================================= // // For case: // try_cast/cast(expr as data_type) op literal Expr::BinaryExpr(BinaryExpr { left, op, right }) - if is_cast_expr_and_support_unwrap_cast_in_comparison_for_binary( - info, &left, op, &right, - ) && op.supports_propagation() => + if supports_cast_predicate_for_binary(info, &left, op, &right) + && op.supports_propagation() => { - unwrap_cast_in_comparison_for_binary(info, *left, *right, op)? + rewrite_cast_predicate_for_binary(info, *left, *right, op)? } // literal op try_cast/cast(expr as data_type) // --> // try_cast/cast(expr as data_type) op_swap literal Expr::BinaryExpr(BinaryExpr { left, op, right }) - if is_cast_expr_and_support_unwrap_cast_in_comparison_for_binary( - info, &right, op, &left, - ) && op.supports_propagation() - && op.swap().is_some() => + if op.supports_propagation() + && op.swap().is_some_and(|swapped| { + supports_cast_predicate_for_binary(info, &right, swapped, &left) + }) => { - unwrap_cast_in_comparison_for_binary( + rewrite_cast_predicate_for_binary( info, *right, *left, @@ -2045,52 +2042,11 @@ impl TreeNodeRewriter for Simplifier<'_> { // For case: // try_cast/cast(expr as left_type) in (expr1,expr2,expr3) Expr::InList(InList { - expr: mut left, + expr, list, negated, - }) if is_cast_expr_and_support_unwrap_cast_in_comparison_for_inlist( - info, &left, &list, - ) => - { - let (Expr::TryCast(TryCast { - expr: left_expr, .. - }) - | Expr::Cast(Cast { - expr: left_expr, .. - })) = left.as_mut() - else { - return internal_err!("Expect cast expr, but got {:?}", left)?; - }; - - let expr_type = info.get_data_type(left_expr)?; - let right_exprs = list - .into_iter() - .map(|right| { - match right { - Expr::Literal(right_lit_value, _) => { - // if the right_lit_value can be casted to the type of internal_left_expr - // we need to unwrap the cast for cast/try_cast expr, and add cast to the literal - let Some(value) = try_cast_literal_to_type(&right_lit_value, &expr_type) else { - internal_err!( - "Can't cast the list expr {:?} to type {}", - right_lit_value, &expr_type - )? - }; - Ok(lit(value)) - } - other_expr => internal_err!( - "Only support literal expr to optimize, but the expr is {:?}", - &other_expr - ), - } - }) - .collect::>>()?; - - Transformed::yes(Expr::InList(InList { - expr: std::mem::take(left_expr), - list: right_exprs, - negated, - })) + }) if supports_cast_predicate_for_inlist(info, &expr, &list) => { + rewrite_cast_predicate_for_inlist(info, *expr, list, negated)? } // ======================================= diff --git a/datafusion/optimizer/src/simplify_expressions/mod.rs b/datafusion/optimizer/src/simplify_expressions/mod.rs index e0b53b79d468c..3812f292e14c3 100644 --- a/datafusion/optimizer/src/simplify_expressions/mod.rs +++ b/datafusion/optimizer/src/simplify_expressions/mod.rs @@ -18,6 +18,7 @@ //! [`SimplifyExpressions`] simplifies expressions in the logical plan, //! [`ExprSimplifier`] simplifies individual `Expr`s. +mod cast_preimage; pub mod expr_simplifier; mod inlist_simplifier; mod linear_aggregates; @@ -27,7 +28,6 @@ pub mod simplify_exprs; pub mod simplify_literal; mod simplify_predicates; mod udf_preimage; -mod unwrap_cast; mod utils; // backwards compatibility diff --git a/datafusion/optimizer/src/simplify_expressions/unwrap_cast.rs b/datafusion/optimizer/src/simplify_expressions/unwrap_cast.rs deleted file mode 100644 index d0165b568a753..0000000000000 --- a/datafusion/optimizer/src/simplify_expressions/unwrap_cast.rs +++ /dev/null @@ -1,712 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -//! Unwrap casts in binary comparisons -//! -//! The functions in this module attempt to remove casts from -//! comparisons to literals ([`ScalarValue`]s) by applying the casts -//! to the literals if possible. It is inspired by the optimizer rule -//! `UnwrapCastInBinaryComparison` of Spark. -//! -//! Removing casts often improves performance because: -//! 1. The cast is done once (to the literal) rather than to every value -//! 2. Can enable other optimizations such as predicate pushdown that -//! don't support casting -//! -//! The rule is applied to expressions of the following forms: -//! -//! 1. `cast(left_expr as data_type) comparison_op literal_expr` -//! 2. `literal_expr comparison_op cast(left_expr as data_type)` -//! 3. `cast(literal_expr) IN (expr1, expr2, ...)` -//! 4. `literal_expr IN (cast(expr1) , cast(expr2), ...)` -//! -//! If the expression matches one of the forms above, the rule will -//! ensure the value of `literal` is in range(min, max) of the -//! expr's data_type, and if the scalar is within range, the literal -//! will be casted to the data type of expr on the other side, and the -//! cast will be removed from the other side. -//! -//! # Example -//! -//! If the DataType of c1 is INT32. Given the filter -//! -//! ```text -//! cast(c1 as INT64) > INT64(10)` -//! ``` -//! -//! This rule will remove the cast and rewrite the expression to: -//! -//! ```text -//! c1 > INT32(10) -//! ``` - -use arrow::datatypes::DataType; -use datafusion_common::{Result, ScalarValue}; -use datafusion_common::{internal_err, tree_node::Transformed}; -use datafusion_expr::{BinaryExpr, lit}; -use datafusion_expr::{Cast, Expr, Operator, TryCast, simplify::SimplifyContext}; -use datafusion_expr_common::casts::{ - is_date_narrowing_cast, is_supported_type, is_timestamp_precision_narrowing_cast, - try_cast_literal_to_type, -}; - -pub(super) fn unwrap_cast_in_comparison_for_binary( - info: &SimplifyContext, - cast_expr: Expr, - literal: Expr, - op: Operator, -) -> Result> { - match (cast_expr, literal) { - ( - Expr::TryCast(TryCast { expr, .. }) | Expr::Cast(Cast { expr, .. }), - Expr::Literal(lit_value, _), - ) => { - let Ok(expr_type) = info.get_data_type(&expr) else { - return internal_err!("Can't get the data type of the expr {:?}", &expr); - }; - - if let Some(value) = cast_literal_to_type_with_op(&lit_value, &expr_type, op) - { - return Ok(Transformed::yes(Expr::BinaryExpr(BinaryExpr { - left: expr, - op, - right: Box::new(lit(value)), - }))); - } - - // if the lit_value can be casted to the type of internal_left_expr - // we need to unwrap the cast for cast/try_cast expr, and add cast to the literal - let Some(value) = try_cast_literal_to_type(&lit_value, &expr_type) else { - return internal_err!( - "Can't cast the literal expr {:?} to type {}", - &lit_value, - &expr_type - ); - }; - Ok(Transformed::yes(Expr::BinaryExpr(BinaryExpr { - left: expr, - op, - right: Box::new(lit(value)), - }))) - } - _ => internal_err!("Expect cast expr and literal"), - } -} - -pub(super) fn is_cast_expr_and_support_unwrap_cast_in_comparison_for_binary( - info: &SimplifyContext, - expr: &Expr, - op: Operator, - literal: &Expr, -) -> bool { - match (expr, literal) { - ( - Expr::TryCast(TryCast { - expr: left_expr, - field, - }) - | Expr::Cast(Cast { - expr: left_expr, - field, - }), - Expr::Literal(lit_val, _), - ) => { - let Ok(expr_type) = info.get_data_type(left_expr) else { - return false; - }; - - let Ok(lit_type) = info.get_data_type(literal) else { - return false; - }; - - if is_timestamp_precision_narrowing_cast(&expr_type, field.data_type()) - || is_date_narrowing_cast(&expr_type, field.data_type()) - { - return false; - } - - if cast_literal_to_type_with_op(lit_val, &expr_type, op).is_some() { - return true; - } - - try_cast_literal_to_type(lit_val, &expr_type).is_some() - && is_supported_type(&expr_type) - && is_supported_type(&lit_type) - } - _ => false, - } -} - -pub(super) fn is_cast_expr_and_support_unwrap_cast_in_comparison_for_inlist( - info: &SimplifyContext, - expr: &Expr, - list: &[Expr], -) -> bool { - let (Expr::TryCast(TryCast { - expr: left_expr, - field, - }) - | Expr::Cast(Cast { - expr: left_expr, - field, - })) = expr - else { - return false; - }; - - let Ok(expr_type) = info.get_data_type(left_expr) else { - return false; - }; - - if !is_supported_type(&expr_type) { - return false; - } - - if is_timestamp_precision_narrowing_cast(&expr_type, field.data_type()) - || is_date_narrowing_cast(&expr_type, field.data_type()) - { - return false; - } - - for right in list { - let Ok(right_type) = info.get_data_type(right) else { - return false; - }; - - if !is_supported_type(&right_type) { - return false; - } - - match right { - Expr::Literal(lit_val, _) - if try_cast_literal_to_type(lit_val, &expr_type).is_some() => {} - _ => return false, - } - } - - true -} - -///// Tries to move a cast from an expression (such as column) to the literal other side of a comparison operator./ -/// -/// Specifically, rewrites -/// ```sql -/// cast(col) -/// ``` -/// -/// To -/// -/// ```sql -/// col cast() -/// col -/// ``` -fn cast_literal_to_type_with_op( - lit_value: &ScalarValue, - target_type: &DataType, - op: Operator, -) -> Option { - match (op, lit_value) { - ( - Operator::Eq | Operator::NotEq, - ScalarValue::Utf8(Some(_)) - | ScalarValue::Utf8View(Some(_)) - | ScalarValue::LargeUtf8(Some(_)), - ) => { - // Only try for integer types (TODO can we do this for other types - // like timestamps)? - use DataType::*; - if matches!( - target_type, - Int8 | Int16 | Int32 | Int64 | UInt8 | UInt16 | UInt32 | UInt64 - ) { - let casted = lit_value.cast_to(target_type).ok()?; - let round_tripped = casted.cast_to(&lit_value.data_type()).ok()?; - if lit_value != &round_tripped { - return None; - } - Some(casted) - } else { - None - } - } - _ => None, - } -} - -#[cfg(test)] -mod tests { - use super::*; - use std::collections::HashMap; - use std::sync::Arc; - - use crate::simplify_expressions::ExprSimplifier; - use arrow::datatypes::{Field, TimeUnit}; - use datafusion_common::{DFSchema, DFSchemaRef}; - use datafusion_expr::simplify::SimplifyContext; - use datafusion_expr::{cast, col, in_list, try_cast}; - - #[test] - fn test_not_unwrap_cast_comparison() { - let schema = expr_test_schema(); - // cast(INT32(c1), INT64) > INT64(c2) - let c1_gt_c2 = cast(col("c1"), DataType::Int64).gt(col("c2")); - assert_eq!(optimize_test(c1_gt_c2.clone(), &schema), c1_gt_c2); - - // INT32(c1) < INT32(16), the type is same - let expr_lt = col("c1").lt(lit(16i32)); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // the 99999999999 is not within the range of MAX(int32) and MIN(int32), we don't cast the lit(99999999999) to int32 type - let expr_lt = cast(col("c1"), DataType::Int64).lt(lit(99999999999i64)); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // cast(c1, UTF8) < '123', only eq/not_eq should be optimized - let expr_lt = cast(col("c1"), DataType::Utf8).lt(lit("123")); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // cast(c1, UTF8) = '0123', cast(cast('0123', Int32), UTF8) != '0123', so '0123' should not - // be casted - let expr_lt = cast(col("c1"), DataType::Utf8).lt(lit("0123")); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // cast(c1, UTF8) = 'not a number', should not be able to cast to column type - let expr_input = cast(col("c1"), DataType::Utf8).eq(lit("not a number")); - assert_eq!(optimize_test(expr_input.clone(), &schema), expr_input); - - // cast(c1, UTF8) = '99999999999', where '99999999999' does not fit into int32, so it will - // not be optimized to integer comparison - let expr_input = cast(col("c1"), DataType::Utf8).eq(lit("99999999999")); - assert_eq!(optimize_test(expr_input.clone(), &schema), expr_input); - } - - #[test] - fn test_unwrap_cast_comparison() { - let schema = expr_test_schema(); - // cast(c1, INT64) < INT64(16) -> INT32(c1) < cast(INT32(16)) - // the 16 is within the range of MAX(int32) and MIN(int32), we can cast the 16 to int32(16) - let expr_lt = cast(col("c1"), DataType::Int64).lt(lit(16i64)); - let expected = col("c1").lt(lit(16i32)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - let expr_lt = try_cast(col("c1"), DataType::Int64).lt(lit(16i64)); - let expected = col("c1").lt(lit(16i32)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // cast(c2, INT32) = INT32(16) => INT64(c2) = INT64(16) - let c2_eq_lit = cast(col("c2"), DataType::Int32).eq(lit(16i32)); - let expected = col("c2").eq(lit(16i64)); - assert_eq!(optimize_test(c2_eq_lit, &schema), expected); - - // cast(c1, INT64) < INT64(NULL) => NULL - let c1_lt_lit_null = cast(col("c1"), DataType::Int64).lt(null_i64()); - let expected = null_bool(); - assert_eq!(optimize_test(c1_lt_lit_null, &schema), expected); - - // cast(INT8(NULL), INT32) < INT32(12) => INT8(NULL) < INT8(12) => BOOL(NULL) - let lit_lt_lit = cast(null_i8(), DataType::Int32).lt(lit(12i32)); - let expected = null_bool(); - assert_eq!(optimize_test(lit_lt_lit, &schema), expected); - - // cast(c1, UTF8) = '123' => c1 = 123 - let expr_input = cast(col("c1"), DataType::Utf8).eq(lit("123")); - let expected = col("c1").eq(lit(123i32)); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // cast(c1, UTF8) != '123' => c1 != 123 - let expr_input = cast(col("c1"), DataType::Utf8).not_eq(lit("123")); - let expected = col("c1").not_eq(lit(123i32)); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // cast(c1, UTF8) = NULL => NULL - let expr_input = cast(col("c1"), DataType::Utf8).eq(lit(ScalarValue::Utf8(None))); - let expected = null_bool(); - assert_eq!(optimize_test(expr_input, &schema), expected); - } - - #[test] - fn test_unwrap_cast_comparison_unsigned() { - // "cast(c6, UINT64) = 0u64 => c6 = 0u32 - let schema = expr_test_schema(); - let expr_input = cast(col("c6"), DataType::UInt64).eq(lit(0u64)); - let expected = col("c6").eq(lit(0u32)); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // cast(c6, UTF8) = "123" => c6 = 123 - let expr_input = cast(col("c6"), DataType::Utf8).eq(lit("123")); - let expected = col("c6").eq(lit(123u32)); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // cast(c6, UTF8) != "123" => c6 != 123 - let expr_input = cast(col("c6"), DataType::Utf8).not_eq(lit("123")); - let expected = col("c6").not_eq(lit(123u32)); - assert_eq!(optimize_test(expr_input, &schema), expected); - } - - #[test] - fn test_unwrap_cast_comparison_string() { - let schema = expr_test_schema(); - let dict = ScalarValue::Dictionary( - Box::new(DataType::Int32), - Box::new(ScalarValue::from("value")), - ); - - // cast(str1 as Dictionary) = arrow_cast('value', 'Dictionary') => str1 = Utf8('value1') - let expr_input = cast(col("str1"), dict.data_type()).eq(lit(dict.clone())); - let expected = col("str1").eq(lit("value")); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // cast(tag as Utf8) = Utf8('value') => tag = arrow_cast('value', 'Dictionary') - let expr_input = cast(col("tag"), DataType::Utf8).eq(lit("value")); - let expected = col("tag").eq(lit(dict.clone())); - assert_eq!(optimize_test(expr_input, &schema), expected); - - // Verify reversed argument order - // arrow_cast('value', 'Dictionary') = cast(str1 as Dictionary) => Utf8('value1') = str1 - let expr_input = lit(dict.clone()).eq(cast(col("str1"), dict.data_type())); - let expected = col("str1").eq(lit("value")); - assert_eq!(optimize_test(expr_input, &schema), expected); - } - - #[test] - fn test_unwrap_cast_comparison_large_string() { - let schema = expr_test_schema(); - // cast(largestr as Dictionary) = arrow_cast('value', 'Dictionary') => str1 = LargeUtf8('value1') - let dict = ScalarValue::Dictionary( - Box::new(DataType::Int32), - Box::new(ScalarValue::LargeUtf8(Some("value".to_owned()))), - ); - let expr_input = cast(col("largestr"), dict.data_type()).eq(lit(dict)); - let expected = - col("largestr").eq(lit(ScalarValue::LargeUtf8(Some("value".to_owned())))); - assert_eq!(optimize_test(expr_input, &schema), expected); - } - - #[test] - fn test_not_unwrap_cast_with_decimal_comparison() { - let schema = expr_test_schema(); - // integer to decimal: value is out of the bounds of the decimal - // cast(c3, INT64) = INT64(100000000000000000) - let expr_eq = cast(col("c3"), DataType::Int64).eq(lit(100000000000000000i64)); - assert_eq!(optimize_test(expr_eq.clone(), &schema), expr_eq); - - // cast(c4, INT64) = INT64(1000) will overflow the i128 - let expr_eq = cast(col("c4"), DataType::Int64).eq(lit(1000i64)); - assert_eq!(optimize_test(expr_eq.clone(), &schema), expr_eq); - - // decimal to decimal: value will lose the scale when convert to the target data type - // c3 = DECIMAL(12340,20,4) - let expr_eq = - cast(col("c3"), DataType::Decimal128(20, 4)).eq(lit_decimal(12340, 20, 4)); - assert_eq!(optimize_test(expr_eq.clone(), &schema), expr_eq); - - // decimal to integer - // c1 = DECIMAL(123, 10, 1): value will lose the scale when convert to the target data type - let expr_eq = - cast(col("c1"), DataType::Decimal128(10, 1)).eq(lit_decimal(123, 10, 1)); - assert_eq!(optimize_test(expr_eq.clone(), &schema), expr_eq); - - // c1 = DECIMAL(1230, 10, 2): value will lose the scale when convert to the target data type - let expr_eq = - cast(col("c1"), DataType::Decimal128(10, 2)).eq(lit_decimal(1230, 10, 2)); - assert_eq!(optimize_test(expr_eq.clone(), &schema), expr_eq); - } - - #[test] - fn test_unwrap_cast_with_decimal_lit_comparison() { - let schema = expr_test_schema(); - // integer to decimal - // c3 < INT64(16) -> c3 < (CAST(INT64(16) AS DECIMAL(18,2)); - let expr_lt = try_cast(col("c3"), DataType::Int64).lt(lit(16i64)); - let expected = col("c3").lt(lit_decimal(1600, 18, 2)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // c3 < INT64(NULL) - let c1_lt_lit_null = cast(col("c3"), DataType::Int64).lt(null_i64()); - let expected = null_bool(); - assert_eq!(optimize_test(c1_lt_lit_null, &schema), expected); - - // decimal to decimal - // c3 < Decimal(123,10,0) -> c3 < CAST(DECIMAL(123,10,0) AS DECIMAL(18,2)) -> c3 < DECIMAL(12300,18,2) - let expr_lt = - cast(col("c3"), DataType::Decimal128(10, 0)).lt(lit_decimal(123, 10, 0)); - let expected = col("c3").lt(lit_decimal(12300, 18, 2)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // c3 < Decimal(1230,10,3) -> c3 < CAST(DECIMAL(1230,10,3) AS DECIMAL(18,2)) -> c3 < DECIMAL(123,18,2) - let expr_lt = - cast(col("c3"), DataType::Decimal128(10, 3)).lt(lit_decimal(1230, 10, 3)); - let expected = col("c3").lt(lit_decimal(123, 18, 2)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // decimal to integer - // c1 < Decimal(12300, 10, 2) -> c1 < CAST(DECIMAL(12300,10,2) AS INT32) -> c1 < INT32(123) - let expr_lt = - cast(col("c1"), DataType::Decimal128(10, 2)).lt(lit_decimal(12300, 10, 2)); - let expected = col("c1").lt(lit(123i32)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - } - - #[test] - fn test_not_unwrap_list_cast_lit_comparison() { - let schema = expr_test_schema(); - // internal left type is not supported - // FLOAT32(C5) in ... - let expr_lt = - cast(col("c5"), DataType::Int64).in_list(vec![lit(12i64), lit(12i64)], false); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // cast(INT32(C1), Float32) in (FLOAT32(1.23), Float32(12), Float32(12)) - let expr_lt = cast(col("c1"), DataType::Float32) - .in_list(vec![lit(12.0f32), lit(12.0f32), lit(1.23f32)], false); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // INT32(C1) in (INT64(99999999999), INT64(12)) - let expr_lt = cast(col("c1"), DataType::Int64) - .in_list(vec![lit(12i32), lit(99999999999i64)], false); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - - // DECIMAL(C3) in (INT64(12), INT32(12), DECIMAL(128,12,3)) - let expr_lt = cast(col("c3"), DataType::Decimal128(12, 3)).in_list( - vec![ - lit_decimal(12, 12, 3), - lit_decimal(12, 12, 3), - lit_decimal(128, 12, 3), - ], - false, - ); - assert_eq!(optimize_test(expr_lt.clone(), &schema), expr_lt); - } - - #[test] - fn test_unwrap_list_cast_comparison() { - let schema = expr_test_schema(); - // INT32(C1) IN (INT32(12),INT64(23),INT64(34),INT64(56),INT64(78)) -> - // INT32(C1) IN (INT32(12),INT32(23),INT32(34),INT32(56),INT32(78)) - let expr_lt = cast(col("c1"), DataType::Int64).in_list( - vec![lit(12i64), lit(23i64), lit(34i64), lit(56i64), lit(78i64)], - false, - ); - let expected = col("c1").in_list( - vec![lit(12i32), lit(23i32), lit(34i32), lit(56i32), lit(78i32)], - false, - ); - assert_eq!(optimize_test(expr_lt, &schema), expected); - // INT32(C2) IN (INT64(NULL),INT64(24),INT64(34),INT64(56),INT64(78)) -> - // INT32(C2) IN (INT32(NULL),INT32(24),INT32(34),INT32(56),INT32(78)) - let expr_lt = cast(col("c2"), DataType::Int32).in_list( - vec![null_i32(), lit(24i32), lit(34i64), lit(56i64), lit(78i64)], - false, - ); - let expected = col("c2").in_list( - vec![null_i64(), lit(24i64), lit(34i64), lit(56i64), lit(78i64)], - false, - ); - - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // decimal test case - // c3 is decimal(18,2) - let expr_lt = cast(col("c3"), DataType::Decimal128(19, 3)).in_list( - vec![ - lit_decimal(12000, 19, 3), - lit_decimal(24000, 19, 3), - lit_decimal(1280, 19, 3), - lit_decimal(1240, 19, 3), - ], - false, - ); - let expected = col("c3").in_list( - vec![ - lit_decimal(1200, 18, 2), - lit_decimal(2400, 18, 2), - lit_decimal(128, 18, 2), - lit_decimal(124, 18, 2), - ], - false, - ); - assert_eq!(optimize_test(expr_lt, &schema), expected); - - // cast(INT32(12), INT64) IN (.....) => - // INT64(12) IN (INT64(12),INT64(13),INT64(14),INT64(15),INT64(16)) - // => true - let expr_lt = cast(lit(12i32), DataType::Int64).in_list( - vec![lit(12i64), lit(13i64), lit(14i64), lit(15i64), lit(16i64)], - false, - ); - let expected = lit(true); - assert_eq!(optimize_test(expr_lt, &schema), expected); - } - - #[test] - fn aliased() { - let schema = expr_test_schema(); - // c1 < INT64(16) -> c1 < cast(INT32(16)) - // the 16 is within the range of MAX(int32) and MIN(int32), we can cast the 16 to int32(16) - let expr_lt = cast(col("c1"), DataType::Int64).lt(lit(16i64)).alias("x"); - let expected = col("c1").lt(lit(16i32)).alias("x"); - assert_eq!(optimize_test(expr_lt, &schema), expected); - } - - #[test] - fn nested() { - let schema = expr_test_schema(); - // c1 < INT64(16) OR c1 > INT64(32) -> c1 < INT32(16) OR c1 > INT32(32) - // the 16 and 32 are within the range of MAX(int32) and MIN(int32), we can cast them to int32 - let expr_lt = cast(col("c1"), DataType::Int64).lt(lit(16i64)).or(cast( - col("c1"), - DataType::Int64, - ) - .gt(lit(32i64))); - let expected = col("c1").lt(lit(16i32)).or(col("c1").gt(lit(32i32))); - assert_eq!(optimize_test(expr_lt, &schema), expected); - } - - #[test] - fn test_not_support_data_type() { - // "c6 > 0" will be cast to `cast(c6 as float) > 0 - // but the type of c6 is uint32 - // the rewriter will not throw error and just return the original expr - let schema = expr_test_schema(); - let expr_input = cast(col("c6"), DataType::Float64).eq(lit(0f64)); - assert_eq!(optimize_test(expr_input.clone(), &schema), expr_input); - - // inlist for unsupported data type - let expr_input = in_list( - cast(col("c6"), DataType::Float64), - // need more literals to avoid rewriting to binary expr - vec![lit(0f64), lit(1f64), lit(2f64), lit(3f64), lit(4f64)], - false, - ); - assert_eq!(optimize_test(expr_input.clone(), &schema), expr_input); - } - - #[test] - /// Basic integration test for unwrapping casts with different timezones - fn test_unwrap_cast_with_timestamp_nanos() { - let schema = expr_test_schema(); - // cast(ts_nano as Timestamp(Nanosecond, UTC)) < 1666612093000000000::Timestamp(Nanosecond, Utc)) - let expr_lt = try_cast(col("ts_nano_none"), timestamp_nano_utc_type()) - .lt(lit_timestamp_nano_utc(1666612093000000000)); - let expected = - col("ts_nano_none").lt(lit_timestamp_nano_none(1666612093000000000)); - assert_eq!(optimize_test(expr_lt, &schema), expected); - } - - #[test] - fn test_not_unwrap_cast_timestamp_precision_narrowing() { - let schema = expr_test_schema(); - let expr_input = cast(col("ts_nano_none"), timestamp_millis_none_type()) - .eq(lit_timestamp_millis_none(1)); - - assert_eq!(optimize_test(expr_input.clone(), &schema), expr_input); - } - - #[test] - fn test_unwrap_cast_timestamp_precision_widening() { - let schema = expr_test_schema(); - let expr_input = cast(col("ts_millis_none"), timestamp_nano_none_type()) - .eq(lit_timestamp_nano_none(1_000_000)); - let expected = col("ts_millis_none").eq(lit_timestamp_millis_none(1)); - - assert_eq!(optimize_test(expr_input, &schema), expected); - } - - fn optimize_test(expr: Expr, schema: &DFSchemaRef) -> Expr { - let simplifier = ExprSimplifier::new( - SimplifyContext::builder() - .with_schema(Arc::clone(schema)) - .build(), - ); - - simplifier.simplify(expr).unwrap() - } - - fn expr_test_schema() -> DFSchemaRef { - Arc::new( - DFSchema::from_unqualified_fields( - vec![ - Field::new("c1", DataType::Int32, false), - Field::new("c2", DataType::Int64, false), - Field::new("c3", DataType::Decimal128(18, 2), false), - Field::new("c4", DataType::Decimal128(38, 37), false), - Field::new("c5", DataType::Float32, false), - Field::new("c6", DataType::UInt32, false), - Field::new("ts_nano_none", timestamp_nano_none_type(), false), - Field::new("ts_millis_none", timestamp_millis_none_type(), false), - Field::new("ts_nano_utf", timestamp_nano_utc_type(), false), - Field::new("str1", DataType::Utf8, false), - Field::new("largestr", DataType::LargeUtf8, false), - Field::new("tag", dictionary_tag_type(), false), - ] - .into(), - HashMap::new(), - ) - .unwrap(), - ) - } - - fn null_bool() -> Expr { - lit(ScalarValue::Boolean(None)) - } - - fn null_i8() -> Expr { - lit(ScalarValue::Int8(None)) - } - - fn null_i32() -> Expr { - lit(ScalarValue::Int32(None)) - } - - fn null_i64() -> Expr { - lit(ScalarValue::Int64(None)) - } - - fn lit_decimal(value: i128, precision: u8, scale: i8) -> Expr { - lit(ScalarValue::Decimal128(Some(value), precision, scale)) - } - - fn lit_timestamp_nano_none(ts: i64) -> Expr { - lit(ScalarValue::TimestampNanosecond(Some(ts), None)) - } - - fn lit_timestamp_millis_none(ts: i64) -> Expr { - lit(ScalarValue::TimestampMillisecond(Some(ts), None)) - } - - fn lit_timestamp_nano_utc(ts: i64) -> Expr { - let utc = Some("+0:00".into()); - lit(ScalarValue::TimestampNanosecond(Some(ts), utc)) - } - - fn timestamp_nano_none_type() -> DataType { - DataType::Timestamp(TimeUnit::Nanosecond, None) - } - - fn timestamp_millis_none_type() -> DataType { - DataType::Timestamp(TimeUnit::Millisecond, None) - } - - // this is the type that now() returns - fn timestamp_nano_utc_type() -> DataType { - let utc = Some("+0:00".into()); - DataType::Timestamp(TimeUnit::Nanosecond, utc) - } - - // a dictionary type for storing string tags - fn dictionary_tag_type() -> DataType { - DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8)) - } -} diff --git a/datafusion/optimizer/tests/optimizer_integration.rs b/datafusion/optimizer/tests/optimizer_integration.rs index 26b48c5e1f352..dff730213a584 100644 --- a/datafusion/optimizer/tests/optimizer_integration.rs +++ b/datafusion/optimizer/tests/optimizer_integration.rs @@ -794,7 +794,7 @@ fn extension_node_does_not_block_projection_pruning() -> Result<()> { Projection: t.a, CAST(t.ts AS Timestamp(ms, "UTC")) AS ts Filter: __common_expr_3 > TimestampMillisecond(1000, Some("UTC")) AND __common_expr_3 < TimestampMillisecond(2000, Some("UTC")) Projection: CAST(t.ts AS Timestamp(ms, "UTC")) AS __common_expr_3, t.a, t.ts - TableScan: t projection=[a, ts], partial_filters=[CAST(t.ts AS Timestamp(ms, "UTC")) > TimestampMillisecond(1000, Some("UTC")), CAST(t.ts AS Timestamp(ms, "UTC")) < TimestampMillisecond(2000, Some("UTC"))] + TableScan: t projection=[a, ts], partial_filters=[t.ts >= TimestampNanosecond(1001000000, None), t.ts < TimestampNanosecond(2000000000, None), CAST(t.ts AS Timestamp(ms, "UTC")) > TimestampMillisecond(1000, Some("UTC")), CAST(t.ts AS Timestamp(ms, "UTC")) < TimestampMillisecond(2000, Some("UTC"))] "#, ); diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index 3e67fc8291a4e..140b925fa3888 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -37,12 +37,13 @@ use arrow::datatypes::{DataType, Schema}; use datafusion_common::{Result, ScalarValue, tree_node::Transformed}; use datafusion_expr::Operator; use datafusion_expr_common::casts::{ - is_date_narrowing_cast, is_timestamp_precision_narrowing_cast, - try_cast_literal_to_type, + CastPredicatePreimage, cast_predicate_preimage, is_date_narrowing_cast, }; use crate::PhysicalExpr; -use crate::expressions::{BinaryExpr, CastExpr, Literal, TryCastExpr, lit}; +use crate::expressions::{ + BinaryExpr, CastExpr, Literal, TryCastExpr, is_not_null, is_null, lit, +}; /// Attempts to unwrap casts in comparison expressions. pub(crate) fn unwrap_cast_in_comparison( @@ -69,8 +70,8 @@ fn try_unwrap_cast_binary( ) && binary.op().supports_propagation() && let Some(unwrapped) = try_unwrap_cast_comparison( Arc::clone(inner_expr), - literal.value(), cast_type, + literal.value(), *binary.op(), schema, )? @@ -88,8 +89,8 @@ fn try_unwrap_cast_binary( && binary.op().supports_propagation() && let Some(unwrapped) = try_unwrap_cast_comparison( Arc::clone(inner_expr), - literal.value(), cast_type, + literal.value(), swapped_op, schema, )? @@ -122,28 +123,84 @@ fn extract_cast_info( /// Try to unwrap a cast in comparison by moving the cast to the literal fn try_unwrap_cast_comparison( inner_expr: Arc, - literal_value: &ScalarValue, cast_type: &DataType, + literal_value: &ScalarValue, op: Operator, schema: &Schema, ) -> Result>> { // Get the data type of the inner expression let inner_type = inner_expr.data_type(schema)?; - if is_timestamp_precision_narrowing_cast(&inner_type, cast_type) - || is_date_narrowing_cast(&inner_type, cast_type) - { + if is_date_narrowing_cast(&inner_type, cast_type) { return Ok(None); } - // Try to cast the literal to the inner expression's type - if let Some(casted_literal) = try_cast_literal_to_type(literal_value, &inner_type) { - let literal_expr = lit(casted_literal); - let binary_expr = BinaryExpr::new(inner_expr, op, literal_expr); - return Ok(Some(Arc::new(binary_expr))); + match cast_predicate_preimage(&inner_type, cast_type, op, literal_value)? { + Some(CastPredicatePreimage::Exact(casted_literal)) => { + let literal_expr = lit(casted_literal); + let binary_expr = BinaryExpr::new(inner_expr, op, literal_expr); + Ok(Some(Arc::new(binary_expr))) + } + Some(CastPredicatePreimage::Range(interval)) => { + rewrite_with_preimage(interval, op, inner_expr).map(Some) + } + None => Ok(None), } +} - Ok(None) +fn rewrite_with_preimage( + interval: datafusion_expr_common::interval_arithmetic::Interval, + op: Operator, + expr: Arc, +) -> Result> { + let (lower, upper) = interval.into_bounds(); + let (lower, upper) = (lit(lower), lit(upper)); + + let rewritten_expr = match op { + Operator::Lt => binary(Arc::clone(&expr), Operator::Lt, lower), + Operator::GtEq => binary(Arc::clone(&expr), Operator::GtEq, lower), + Operator::Gt => binary(Arc::clone(&expr), Operator::GtEq, upper), + Operator::LtEq => binary(Arc::clone(&expr), Operator::Lt, upper), + Operator::Eq => binary( + binary(Arc::clone(&expr), Operator::GtEq, lower), + Operator::And, + binary(expr, Operator::Lt, upper), + ), + Operator::NotEq => binary( + binary(Arc::clone(&expr), Operator::Lt, lower), + Operator::Or, + binary(expr, Operator::GtEq, upper), + ), + Operator::IsNotDistinctFrom => binary( + binary( + is_not_null(Arc::clone(&expr))?, + Operator::And, + binary(Arc::clone(&expr), Operator::GtEq, lower), + ), + Operator::And, + binary(expr, Operator::Lt, upper), + ), + Operator::IsDistinctFrom => binary( + binary( + binary(Arc::clone(&expr), Operator::Lt, lower), + Operator::Or, + binary(Arc::clone(&expr), Operator::GtEq, upper), + ), + Operator::Or, + is_null(expr)?, + ), + _ => unreachable!("preimage only supports comparison operators"), + }; + + Ok(rewritten_expr) +} + +fn binary( + left: Arc, + op: Operator, + right: Arc, +) -> Arc { + Arc::new(BinaryExpr::new(left, op, right)) } #[cfg(test)] @@ -578,7 +635,7 @@ mod tests { } #[test] - fn test_not_unwrap_timestamp_precision_narrowing() { + fn test_timestamp_precision_narrowing_range_preimage_gt() { let schema = Schema::new(vec![Field::new( "ts", DataType::Timestamp(TimeUnit::Nanosecond, None), @@ -591,42 +648,61 @@ mod tests { DataType::Timestamp(TimeUnit::Millisecond, None), None, )); - let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1), None)); + let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr)); + Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr)); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); - assert!(!result.transformed); + assert!(result.transformed); + let optimized_binary = result.data.downcast_ref::().unwrap(); + assert_eq!(*optimized_binary.op(), Operator::GtEq); + assert!(!is_cast_expr(optimized_binary.left())); + let right_literal = optimized_binary.right().downcast_ref::().unwrap(); + assert_eq!( + right_literal.value(), + &ScalarValue::TimestampNanosecond(Some(1_001_000_000), None) + ); } #[test] - fn test_unwrap_timestamp_precision_widening() { + fn test_timestamp_precision_narrowing_range_preimage_eq() { let schema = Schema::new(vec![Field::new( "ts", - DataType::Timestamp(TimeUnit::Millisecond, None), + DataType::Timestamp(TimeUnit::Nanosecond, None), false, )]); let column_expr = col("ts", &schema).unwrap(); let cast_expr = Arc::new(CastExpr::new( column_expr, - DataType::Timestamp(TimeUnit::Nanosecond, None), + DataType::Timestamp(TimeUnit::Millisecond, None), None, )); - let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(1_000_000), None)); + let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(-1), None)); let binary_expr = Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr)); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); assert!(result.transformed); - let optimized_binary = result.data.downcast_ref::().unwrap(); - assert!(!is_cast_expr(optimized_binary.left())); - let right_literal = optimized_binary.right().downcast_ref::().unwrap(); + let and_binary = result.data.downcast_ref::().unwrap(); + assert_eq!(*and_binary.op(), Operator::And); + + let lower_binary = and_binary.left().downcast_ref::().unwrap(); + assert_eq!(*lower_binary.op(), Operator::GtEq); + let lower_literal = lower_binary.right().downcast_ref::().unwrap(); assert_eq!( - right_literal.value(), - &ScalarValue::TimestampMillisecond(Some(1), None) + lower_literal.value(), + &ScalarValue::TimestampNanosecond(Some(-1_999_999), None) + ); + + let upper_binary = and_binary.right().downcast_ref::().unwrap(); + assert_eq!(*upper_binary.op(), Operator::Lt); + let upper_literal = upper_binary.right().downcast_ref::().unwrap(); + assert_eq!( + upper_literal.value(), + &ScalarValue::TimestampNanosecond(Some(-999_999), None) ); } From b591b73ddae69a0c609aa53ee8ead49c696dbc7b Mon Sep 17 00:00:00 2001 From: discord9 Date: Thu, 11 Jun 2026 20:29:18 +0800 Subject: [PATCH 02/14] test: more tests Signed-off-by: discord9 --- datafusion/expr-common/src/casts.rs | 82 ++++++++++++++++++- .../src/simplify_expressions/cast_preimage.rs | 78 +++++++++++++++++- .../src/simplifier/unwrap_cast.rs | 44 ++++++++++ .../sqllogictest/test_files/simplify_expr.slt | 55 ++++++++++++- 4 files changed, 256 insertions(+), 3 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 6f69aa15429b9..5be6b16616787 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -137,7 +137,15 @@ pub fn cast_predicate_exact_literal( return None; } - try_cast_literal_to_type(lit_value, source_type) + let source_value = try_cast_literal_to_type(lit_value, source_type)?; + if is_timestamp_cast(source_type, target_type) { + let round_tripped = try_cast_literal_to_type(&source_value, target_type)?; + if &round_tripped != lit_value { + return None; + } + } + + Some(source_value) } /// Returns true when casting a timestamp from `source_type` to `target_type` @@ -155,6 +163,13 @@ pub fn is_timestamp_precision_narrowing_cast( timestamp_unit_scale(source_unit) > timestamp_unit_scale(target_unit) } +fn is_timestamp_cast(source_type: &DataType, target_type: &DataType) -> bool { + matches!( + (source_type, target_type), + (DataType::Timestamp(_, _), DataType::Timestamp(_, _)) + ) +} + fn exact_preimage_for_cast_predicate( source_type: &DataType, target_type: &DataType, @@ -1520,6 +1535,23 @@ mod tests { )) ); + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(0), None), + ) + .unwrap(), + Some(CastPredicatePreimage::Range( + Interval::try_new( + ScalarValue::TimestampNanosecond(Some(-999_999), None), + ScalarValue::TimestampNanosecond(Some(1_000_000), None), + ) + .unwrap() + )) + ); + assert_eq!( cast_predicate_preimage( &ts_ns, @@ -1538,6 +1570,54 @@ mod tests { ); } + #[test] + fn test_cast_predicate_preimage_timestamp_widening_exact_only() { + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + + assert_eq!( + cast_predicate_preimage( + &ts_ms, + &ts_ns, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_000_000), None), + ) + .unwrap(), + Some(CastPredicatePreimage::Exact( + ScalarValue::TimestampMillisecond(Some(123), None) + )) + ); + + assert_eq!( + cast_predicate_preimage( + &ts_ms, + &ts_ns, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_456_789), None), + ) + .unwrap(), + None + ); + } + + #[test] + fn test_cast_predicate_preimage_timestamp_null_literal_unsupported() { + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + let null_ms = ScalarValue::TimestampMillisecond(None, None); + + for op in [ + Operator::Eq, + Operator::IsDistinctFrom, + Operator::IsNotDistinctFrom, + ] { + assert_eq!( + cast_predicate_preimage(&ts_ns, &ts_ms, op, &null_ms).unwrap(), + None + ); + } + } + #[test] fn test_try_cast_to_string_type() { let scalars = vec![ diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index aaf3efb5e612b..7e468da26d40b 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -192,7 +192,7 @@ mod tests { use arrow::datatypes::{Field, TimeUnit}; use datafusion_common::{DFSchema, DFSchemaRef, ScalarValue}; use datafusion_expr::simplify::SimplifyContext; - use datafusion_expr::{cast, col, in_list}; + use datafusion_expr::{binary_expr, cast, col, in_list}; #[test] fn test_cast_predicate_exact_literal_unwrap() { @@ -291,6 +291,81 @@ mod tests { cast(col("ts_nano"), timestamp_millis_type()).lt_eq(lit_timestamp_millis(-1)); let expected = col("ts_nano").lt(lit_timestamp_nano(-999_999)); assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).lt(lit_timestamp_millis(0)); + let expected = col("ts_nano").lt(lit_timestamp_nano(-999_999)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).lt_eq(lit_timestamp_millis(0)); + let expected = col("ts_nano").lt(lit_timestamp_nano(1_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).gt(lit_timestamp_millis(0)); + let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).gt_eq(lit_timestamp_millis(0)); + let expected = col("ts_nano").gt_eq(lit_timestamp_nano(-999_999)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = + cast(col("ts_nano"), timestamp_millis_type()).not_eq(lit_timestamp_millis(0)); + let expected = col("ts_nano") + .lt(lit_timestamp_nano(-999_999)) + .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_cast_preimage_timestamp_precision_narrowing_distinctness() { + let schema = expr_test_schema(); + + let expr = binary_expr( + cast(col("ts_nano"), timestamp_millis_type()), + Operator::IsNotDistinctFrom, + lit_timestamp_millis(0), + ); + let expected = col("ts_nano") + .gt_eq(lit_timestamp_nano(-999_999)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = binary_expr( + cast(col("ts_nano"), timestamp_millis_type()), + Operator::IsDistinctFrom, + lit_timestamp_millis(0), + ); + let expected = col("ts_nano") + .lt(lit_timestamp_nano(-999_999)) + .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_cast_preimage_timestamp_widening_requires_exact_literal() { + let schema = expr_test_schema(); + + let expr = cast(col("ts_milli"), timestamp_nano_type()) + .eq(lit_timestamp_nano(123_000_000)); + let expected = col("ts_milli").eq(lit_timestamp_millis(123)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = cast(col("ts_milli"), timestamp_nano_type()) + .eq(lit_timestamp_nano(123_456_789)); + assert_eq!(optimize_test(expr.clone(), &schema), expr); + + let expr = cast(col("ts_milli"), timestamp_nano_type()) + .gt_eq(lit_timestamp_nano(123_000_000)); + let expected = col("ts_milli").gt_eq(lit_timestamp_millis(123)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = cast(col("ts_milli"), timestamp_nano_type()) + .gt_eq(lit_timestamp_nano(123_456_789)); + assert_eq!(optimize_test(expr.clone(), &schema), expr); } fn optimize_test(expr: Expr, schema: &DFSchemaRef) -> Expr { @@ -308,6 +383,7 @@ mod tests { DFSchema::from_unqualified_fields( vec![ Field::new("c1", DataType::Int32, false), + Field::new("ts_milli", timestamp_millis_type(), false), Field::new("ts_nano", timestamp_nano_type(), false), ] .into(), diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index 140b925fa3888..6993db71e9a93 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -706,6 +706,50 @@ mod tests { ); } + #[test] + fn test_timestamp_widening_exactness() { + let schema = Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(TimeUnit::Millisecond, None), + false, + )]); + + let column_expr = col("ts", &schema).unwrap(); + let cast_expr = Arc::new(CastExpr::new( + column_expr, + DataType::Timestamp(TimeUnit::Nanosecond, None), + None, + )); + let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(123_000_000), None)); + let binary_expr = + Arc::new(BinaryExpr::new(cast_expr, Operator::GtEq, literal_expr)); + + let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + + assert!(result.transformed); + let optimized_binary = result.data.downcast_ref::().unwrap(); + assert_eq!(*optimized_binary.op(), Operator::GtEq); + assert!(!is_cast_expr(optimized_binary.left())); + let right_literal = optimized_binary.right().downcast_ref::().unwrap(); + assert_eq!( + right_literal.value(), + &ScalarValue::TimestampMillisecond(Some(123), None) + ); + + let column_expr = col("ts", &schema).unwrap(); + let cast_expr = Arc::new(CastExpr::new( + column_expr, + DataType::Timestamp(TimeUnit::Nanosecond, None), + None, + )); + let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(123_456_789), None)); + let binary_expr = + Arc::new(BinaryExpr::new(cast_expr, Operator::GtEq, literal_expr)); + + let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + assert!(!result.transformed); + } + #[test] fn test_complex_nested_expression() { let schema = test_schema(); diff --git a/datafusion/sqllogictest/test_files/simplify_expr.slt b/datafusion/sqllogictest/test_files/simplify_expr.slt index c70ff2b955e19..8da390397867c 100644 --- a/datafusion/sqllogictest/test_files/simplify_expr.slt +++ b/datafusion/sqllogictest/test_files/simplify_expr.slt @@ -499,7 +499,6 @@ select id from date_unwrap where cast(d64 as date) = DATE '1970-01-01' order by statement ok drop table date_unwrap; - # IN-list and array_has simplifications must preserve SQL NULL semantics. query BBB SELECT @@ -518,3 +517,57 @@ SELECT FROM (VALUES (2, NULL::INT)) AS t(i, n); ---- NULL NULL + +# Cast predicate preimage rewrites for timestamp precision casts +statement ok +CREATE TABLE cast_preimage_ts AS +SELECT + arrow_cast(0, 'Timestamp(ns)') AS ts_nano, + arrow_cast(0, 'Timestamp(ms)') AS ts_milli; + +query TT +EXPLAIN SELECT * FROM cast_preimage_ts +WHERE CAST(ts_nano AS TIMESTAMP(3)) > arrow_cast(1000, 'Timestamp(ms)'); +---- +logical_plan +01)Filter: cast_preimage_ts.ts_nano >= TimestampNanosecond(1001000000, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: ts_nano@0 >= 1001000000 +02)--DataSourceExec: partitions=1, partition_sizes=[1] + +query TT +EXPLAIN SELECT * FROM cast_preimage_ts +WHERE CAST(ts_nano AS TIMESTAMP(3)) IS DISTINCT FROM arrow_cast(0, 'Timestamp(ms)'); +---- +logical_plan +01)Filter: cast_preimage_ts.ts_nano < TimestampNanosecond(-999999, None) OR cast_preimage_ts.ts_nano >= TimestampNanosecond(1000000, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: ts_nano@0 < -999999 OR ts_nano@0 >= 1000000 +02)--DataSourceExec: partitions=1, partition_sizes=[1] + +query TT +EXPLAIN SELECT * FROM cast_preimage_ts +WHERE CAST(ts_milli AS TIMESTAMP(9)) >= arrow_cast(123000000, 'Timestamp(ns)'); +---- +logical_plan +01)Filter: cast_preimage_ts.ts_milli >= TimestampMillisecond(123, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: ts_milli@1 >= 123 +02)--DataSourceExec: partitions=1, partition_sizes=[1] + +query TT +EXPLAIN SELECT * FROM cast_preimage_ts +WHERE CAST(ts_milli AS TIMESTAMP(9)) >= arrow_cast(123456789, 'Timestamp(ns)'); +---- +logical_plan +01)Filter: CAST(cast_preimage_ts.ts_milli AS Timestamp(ns)) >= TimestampNanosecond(123456789, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: CAST(ts_milli@1 AS Timestamp(ns)) >= 123456789 +02)--DataSourceExec: partitions=1, partition_sizes=[1] + +statement ok +DROP TABLE cast_preimage_ts; From d801a43a3c7161a3b077051e1a19f0dfb117398e Mon Sep 17 00:00:00 2001 From: discord9 Date: Fri, 12 Jun 2026 11:09:42 +0800 Subject: [PATCH 03/14] chore: clippy Signed-off-by: discord9 --- .../optimizer/src/simplify_expressions/cast_preimage.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index 7e468da26d40b..1fb17fb94462f 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -130,7 +130,7 @@ pub(super) fn rewrite_cast_predicate_for_inlist( list: Vec, negated: bool, ) -> Result> { - let Some((inner_expr, _target_type)) = cast_input_and_type(expr) else { + let Some((inner_expr, target_type)) = cast_input_and_type(expr) else { return internal_err!("Expect cast expr"); }; let source_type = info.get_data_type(&inner_expr)?; @@ -140,7 +140,7 @@ pub(super) fn rewrite_cast_predicate_for_inlist( .map(|right| match right { Expr::Literal(lit_value, _) => { let Some(value) = - cast_predicate_exact_literal(&source_type, &_target_type, &lit_value) + cast_predicate_exact_literal(&source_type, &target_type, &lit_value) else { return internal_err!( "Can't cast the list expr {:?} to type {}", From 641fd8cc092ba0b21a105cd6c082154e08c07083 Mon Sep 17 00:00:00 2001 From: discord9 Date: Mon, 15 Jun 2026 14:24:14 +0800 Subject: [PATCH 04/14] test: cover cast preimage distinctness --- datafusion/expr-common/src/casts.rs | 8 ++ .../src/simplifier/unwrap_cast.rs | 120 ++++++++++++++++++ 2 files changed, 128 insertions(+) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 5be6b16616787..1550660c7f8d1 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -228,6 +228,14 @@ fn cast_to_string_equality_preimage( } } +/// Computes the source-domain preimage interval for timestamp precision +/// narrowing casts. +/// +/// The preimage is computed entirely in the source timestamp domain and +/// preserves `source_tz`. The target timezone is intentionally *not* copied +/// into the generated bounds: this models the raw timestamp values after +/// truncation and avoids comparing source columns against target-timezone +/// literals. fn timestamp_precision_narrowing_preimage( source_type: &DataType, target_type: &DataType, diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index 6993db71e9a93..4b6457422ff46 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -750,6 +750,126 @@ mod tests { assert!(!result.transformed); } + #[test] + fn test_timestamp_precision_narrowing_range_preimage_is_distinct_from() { + let schema = Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + )]); + + let column_expr = col("ts", &schema).unwrap(); + let cast_expr = Arc::new(CastExpr::new( + column_expr, + DataType::Timestamp(TimeUnit::Millisecond, None), + None, + )); + let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); + let binary_expr = Arc::new(BinaryExpr::new( + cast_expr, + Operator::IsDistinctFrom, + literal_expr, + )); + + let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + + assert!(result.transformed); + + // Expected: OR( OR(expr < lower, expr >= upper), IS NULL(expr) ) + let outer_or = result.data.downcast_ref::().unwrap(); + assert_eq!(*outer_or.op(), Operator::Or); + + // Right side of outer OR → IS NULL + assert!( + outer_or + .right() + .downcast_ref::() + .is_some() + ); + + // Left side of outer OR → OR(expr < lower, expr >= upper) + let inner_or = outer_or.left().downcast_ref::().unwrap(); + assert_eq!(*inner_or.op(), Operator::Or); + + // Left-left: expr < lower + let lt_binary = inner_or.left().downcast_ref::().unwrap(); + assert_eq!(*lt_binary.op(), Operator::Lt); + let lt_literal = lt_binary.right().downcast_ref::().unwrap(); + assert_eq!( + lt_literal.value(), + &ScalarValue::TimestampNanosecond(Some(1_000_000_000), None) + ); + + // Left-right: expr >= upper + let gte_binary = inner_or.right().downcast_ref::().unwrap(); + assert_eq!(*gte_binary.op(), Operator::GtEq); + let gte_literal = gte_binary.right().downcast_ref::().unwrap(); + assert_eq!( + gte_literal.value(), + &ScalarValue::TimestampNanosecond(Some(1_001_000_000), None) + ); + } + + #[test] + fn test_timestamp_precision_narrowing_range_preimage_is_not_distinct_from() { + let schema = Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + )]); + + let column_expr = col("ts", &schema).unwrap(); + let cast_expr = Arc::new(CastExpr::new( + column_expr, + DataType::Timestamp(TimeUnit::Millisecond, None), + None, + )); + let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); + let binary_expr = Arc::new(BinaryExpr::new( + cast_expr, + Operator::IsNotDistinctFrom, + literal_expr, + )); + + let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + + assert!(result.transformed); + + // Expected: AND( AND(IS NOT NULL(expr), expr >= lower), expr < upper ) + let outer_and = result.data.downcast_ref::().unwrap(); + assert_eq!(*outer_and.op(), Operator::And); + + // Right side of outer AND → expr < upper + let upper_binary = outer_and.right().downcast_ref::().unwrap(); + assert_eq!(*upper_binary.op(), Operator::Lt); + let upper_literal = upper_binary.right().downcast_ref::().unwrap(); + assert_eq!( + upper_literal.value(), + &ScalarValue::TimestampNanosecond(Some(1_001_000_000), None) + ); + + // Left side of outer AND → AND(IS NOT NULL(expr), expr >= lower) + let inner_and = outer_and.left().downcast_ref::().unwrap(); + assert_eq!(*inner_and.op(), Operator::And); + + // Left-left: IS NOT NULL + assert!( + inner_and + .left() + .downcast_ref::() + .is_some() + ); + + // Left-right: expr >= lower + let gte_binary = inner_and.right().downcast_ref::().unwrap(); + assert_eq!(*gte_binary.op(), Operator::GtEq); + let gte_literal = gte_binary.right().downcast_ref::().unwrap(); + assert_eq!( + gte_literal.value(), + &ScalarValue::TimestampNanosecond(Some(1_000_000_000), None) + ); + } + #[test] fn test_complex_nested_expression() { let schema = test_schema(); From 6a6d53a4eb25ff0d8fe0c50aa6ad0655740af066 Mon Sep 17 00:00:00 2001 From: discord9 Date: Mon, 15 Jun 2026 15:17:19 +0800 Subject: [PATCH 05/14] test: cover cast preimage edge cases --- datafusion/expr-common/src/casts.rs | 66 +++++++++++++++++++ .../src/simplify_expressions/cast_preimage.rs | 35 ++++++++++ .../sqllogictest/test_files/simplify_expr.slt | 11 ++++ 3 files changed, 112 insertions(+) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 1550660c7f8d1..8234b09b57aca 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -2061,4 +2061,70 @@ mod tests { ExpectedCast::NoValue, ); } + + #[test] + fn test_cast_predicate_preimage_extreme_literals() { + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + + // i64::MAX in milliseconds expands beyond i64 range in nanoseconds, + // so the preimage should return None rather than panicking. + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(i64::MAX), None), + ) + .unwrap(), + None + ); + + // i64::MIN in milliseconds expands beyond i64 range in nanoseconds. + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(i64::MIN), None), + ) + .unwrap(), + None + ); + } + + #[test] + fn test_cast_predicate_preimage_timezone_preservation() { + let source_type = + DataType::Timestamp(TimeUnit::Nanosecond, Some("+05:30".into())); + let target_type = DataType::Timestamp(TimeUnit::Millisecond, Some("UTC".into())); + let lit = ScalarValue::TimestampMillisecond(Some(1000), Some("UTC".into())); + + let result = + cast_predicate_preimage(&source_type, &target_type, Operator::Eq, &lit) + .unwrap(); + + match result { + Some(CastPredicatePreimage::Range(interval)) => { + let (lower, upper) = interval.into_bounds(); + assert_eq!( + lower, + ScalarValue::TimestampNanosecond( + Some(1_000_000_000), + Some("+05:30".into()) + ), + "lower bound should preserve source timezone +05:30" + ); + assert_eq!( + upper, + ScalarValue::TimestampNanosecond( + Some(1_001_000_000), + Some("+05:30".into()) + ), + "upper bound should preserve source timezone +05:30" + ); + } + other => panic!("Expected CastPredicatePreimage::Range but got {other:?}"), + } + } } diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index 1fb17fb94462f..b03fc44abb45a 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -368,6 +368,41 @@ mod tests { assert_eq!(optimize_test(expr.clone(), &schema), expr); } + #[test] + fn test_cast_preimage_timestamp_literal_left_range() { + let schema = expr_test_schema(); + + // lit_timestamp_millis(1000) < cast(col("ts_nano"), timestamp_millis_type()) + // should swap and rewrite to ts_nano >= 1_001_000_000ns + let expr = + lit_timestamp_millis(1000).lt(cast(col("ts_nano"), timestamp_millis_type())); + let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + // lit_timestamp_millis(1000) <= cast(col("ts_nano"), timestamp_millis_type()) + // should swap and rewrite to ts_nano >= 1_000_000_000ns + let expr = lit_timestamp_millis(1000) + .lt_eq(cast(col("ts_nano"), timestamp_millis_type())); + let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + // lit_timestamp_millis(1000) > cast(col("ts_nano"), timestamp_millis_type()) + // should swap and rewrite to ts_nano < 1_000_000_000ns + let expr = + lit_timestamp_millis(1000).gt(cast(col("ts_nano"), timestamp_millis_type())); + let expected = col("ts_nano").lt(lit_timestamp_nano(1_000_000_000)); + assert_eq!(optimize_test(expr, &schema), expected); + + // lit_timestamp_millis(1000) = cast(col("ts_nano"), timestamp_millis_type()) + // should swap and rewrite to range preimage + let expr = + lit_timestamp_millis(1000).eq(cast(col("ts_nano"), timestamp_millis_type())); + let expected = col("ts_nano") + .gt_eq(lit_timestamp_nano(1_000_000_000)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_001_000_000))); + assert_eq!(optimize_test(expr, &schema), expected); + } + fn optimize_test(expr: Expr, schema: &DFSchemaRef) -> Expr { let simplifier = ExprSimplifier::new( SimplifyContext::builder() diff --git a/datafusion/sqllogictest/test_files/simplify_expr.slt b/datafusion/sqllogictest/test_files/simplify_expr.slt index 8da390397867c..76136596aef95 100644 --- a/datafusion/sqllogictest/test_files/simplify_expr.slt +++ b/datafusion/sqllogictest/test_files/simplify_expr.slt @@ -569,5 +569,16 @@ physical_plan 01)FilterExec: CAST(ts_milli@1 AS Timestamp(ns)) >= 123456789 02)--DataSourceExec: partitions=1, partition_sizes=[1] +query TT +EXPLAIN SELECT * FROM cast_preimage_ts +WHERE CAST(ts_nano AS TIMESTAMP(3)) = arrow_cast(-1, 'Timestamp(ms)'); +---- +logical_plan +01)Filter: cast_preimage_ts.ts_nano >= TimestampNanosecond(-1999999, None) AND cast_preimage_ts.ts_nano < TimestampNanosecond(-999999, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: ts_nano@0 >= -1999999 AND ts_nano@0 < -999999 +02)--DataSourceExec: partitions=1, partition_sizes=[1] + statement ok DROP TABLE cast_preimage_ts; From cd2b8596c8e84169d199d892ef7f799d805fece0 Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 17 Jun 2026 17:34:46 +0800 Subject: [PATCH 06/14] fix: safe check for types Signed-off-by: discord9 --- datafusion/expr-common/src/casts.rs | 453 ++++++++++++++++-- .../src/simplify_expressions/cast_preimage.rs | 52 +- .../physical-expr/src/simplifier/mod.rs | 28 +- .../src/simplifier/unwrap_cast.rs | 15 +- 4 files changed, 493 insertions(+), 55 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 8234b09b57aca..0aaf827eac1a6 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -104,21 +104,36 @@ pub fn cast_predicate_preimage( op: Operator, lit_value: &ScalarValue, ) -> Result> { - if let Some(interval) = - timestamp_precision_narrowing_preimage(source_type, target_type, lit_value)? - { - return Ok(Some(CastPredicatePreimage::Range(interval))); + if let Some(preimage) = maybe_range_preimage(source_type, target_type, lit_value)? { + return Ok(Some(preimage)); } - if is_timestamp_precision_narrowing_cast(source_type, target_type) - || is_date_narrowing_cast(source_type, target_type) + if is_date_narrowing_cast(source_type, target_type) { + return Ok(None); + } + + if let Some(value) = + exact_preimage_int_to_str_eq_like(source_type, target_type, op, lit_value) { + return Ok(Some(CastPredicatePreimage::Exact(value))); + } + + Ok(exact_preimage_cast(source_type, target_type, lit_value) + .map(CastPredicatePreimage::Exact)) +} + +fn maybe_range_preimage( + source_type: &DataType, + target_type: &DataType, + lit_value: &ScalarValue, +) -> Result> { + if !is_timestamp_precision_narrowing_cast(source_type, target_type) { return Ok(None); } Ok( - exact_preimage_for_cast_predicate(source_type, target_type, op, lit_value) - .map(CastPredicatePreimage::Exact), + timestamp_narrowing_range_preimage(source_type, target_type, lit_value)? + .map(CastPredicatePreimage::Range), ) } @@ -126,14 +141,21 @@ pub fn cast_predicate_preimage( /// /// This intentionally returns `None` for timestamp precision narrowing: those /// casts are many-to-one and need range preimages instead. -pub fn cast_predicate_exact_literal( +pub fn exact_preimage_cast( source_type: &DataType, target_type: &DataType, lit_value: &ScalarValue, ) -> Option { - if is_timestamp_precision_narrowing_cast(source_type, target_type) - || is_date_narrowing_cast(source_type, target_type) - { + if is_date_narrowing_cast(source_type, target_type) { + return None; + } + + // Apply a family-level safety gate: the source→target cast must be + // value-preserving over the full source domain. Timestamp precision + // narrowing is handled as a Range preimage, not an Exact preimage; + // timestamp widening / equal-unit casts still require the literal + // round-trip check below. + if !is_exact_cast_safe(source_type, target_type) { return None; } @@ -170,24 +192,183 @@ fn is_timestamp_cast(source_type: &DataType, target_type: &DataType) -> bool { ) } -fn exact_preimage_for_cast_predicate( - source_type: &DataType, - target_type: &DataType, - op: Operator, - lit_value: &ScalarValue, -) -> Option { - cast_to_string_equality_preimage(source_type, target_type, op, lit_value) - .or_else(|| cast_predicate_exact_literal(source_type, target_type, lit_value)) +/// Returns `true` when the cast from `source_type` to `target_type` is +/// value-preserving (injective) and order-preserving at the family level — +/// i.e. every value in the source domain can be round-tripped through the cast +/// without information loss, and comparisons keep the same ordering. +/// +/// This is a conservative gate used by [`exact_preimage_cast`]. Timestamp +/// precision narrowing is not exact-safe (it is handled by +/// [`CastPredicatePreimage::Range`]); timestamp widening / equal-unit casts are +/// family-safe but still require the caller's literal round-trip check. +fn is_exact_cast_safe(source_type: &DataType, target_type: &DataType) -> bool { + // Unwrap at most one level of dictionary for family-level checking. + let (src, tgt) = match (source_type, target_type) { + (DataType::Dictionary(_, v_src), DataType::Dictionary(_, v_tgt)) => { + (v_src.as_ref(), v_tgt.as_ref()) + } + (DataType::Dictionary(_, v_src), tgt) => (v_src.as_ref(), tgt), + (src, DataType::Dictionary(_, v_tgt)) => (src, v_tgt.as_ref()), + _ => (source_type, target_type), + }; + + if src == tgt { + return true; + } + if is_timestamp_cast(src, tgt) { + return !is_timestamp_precision_narrowing_cast(src, tgt); + } + + // Whitelist of family-level safe casts. Anything not listed here is + // considered potentially information-losing and blocked. + if is_integer_type(src) && is_integer_type(tgt) { + return is_safe_integer_cast(src, tgt); + } + if matches!( + (src, tgt), + (DataType::Date32, DataType::Int32) + | (DataType::Int32, DataType::Date32) + | (DataType::Date64, DataType::Int64) + | (DataType::Int64, DataType::Date64) + ) { + // Date32 ↔ Int32 and Date64 ↔ Int64 share the same internal + // representation. + return true; + } + if matches!((src, tgt), (DataType::Date32, DataType::Date64)) { + // Date32 → Date64 is injective: every source day maps to its midnight + // millisecond value. The literal conversion below enforces the exact + // day-boundary requirement in the reverse direction. + return true; + } + if is_supported_string_type(src) && is_supported_string_type(tgt) { + // String width casts among Utf8 / LargeUtf8 / Utf8View. + return true; + } + if is_decimal_type(src) && is_decimal_type(tgt) { + // Decimal widening only: target integer digits and scale must + // be at least as large as source. + return is_safe_decimal_widening(src, tgt); + } + if is_integer_type(src) && is_decimal_type(tgt) { + // int / uint → decimal: allow only when the target can represent + // the full source integer domain. + return is_safe_integer_to_decimal(src, tgt); + } + if matches!((src, tgt), (DataType::FixedSizeBinary(_), DataType::Binary)) { + // FixedSizeBinary(n) → Binary is always widening. + return true; + } + if matches!( + (src, tgt), + (DataType::FixedSizeBinary(n1), DataType::FixedSizeBinary(n2)) if n1 == n2 + ) { + return true; + } + false } -/// Computes a singleton preimage for equality predicates over casts whose target -/// value is a string representation of an integer source value. +/// Returns `true` for integer types (signed and unsigned, excluding Date and +/// Decimal). +fn is_integer_type(dt: &DataType) -> bool { + matches!( + dt, + DataType::Int8 + | DataType::Int16 + | DataType::Int32 + | DataType::Int64 + | DataType::UInt8 + | DataType::UInt16 + | DataType::UInt32 + | DataType::UInt64 + ) +} + +/// Integer cast that is value-preserving over the full source domain. +fn is_safe_integer_cast(src: &DataType, tgt: &DataType) -> bool { + use DataType::*; + match (src, tgt) { + // Signed → wider signed + (Int8, Int16 | Int32 | Int64) => true, + (Int16, Int32 | Int64) => true, + (Int32, Int64) => true, + // Unsigned → wider unsigned + (UInt8, UInt16 | UInt32 | UInt64) => true, + (UInt16, UInt32 | UInt64) => true, + (UInt32, UInt64) => true, + // Unsigned → wider signed (full source range fits in target) + (UInt8, Int16 | Int32 | Int64) => true, + (UInt16, Int32 | Int64) => true, + (UInt32, Int64) => true, + // Anything else is narrowing, partial, or crosses signedness + // without guaranteed domain containment. + _ => false, + } +} + +fn is_decimal_type(dt: &DataType) -> bool { + matches!( + dt, + DataType::Decimal128(_, _) + | DataType::Decimal64(_, _) + | DataType::Decimal32(_, _) + ) +} + +fn decimal_precision_scale(dt: &DataType) -> (u8, i8) { + match dt { + DataType::Decimal128(p, s) + | DataType::Decimal64(p, s) + | DataType::Decimal32(p, s) => (*p, *s), + _ => unreachable!(), + } +} + +/// Target must have at least as many integer digits and at least as large a +/// scale as source. +fn is_safe_decimal_widening(src: &DataType, tgt: &DataType) -> bool { + let (p_src, s_src) = decimal_precision_scale(src); + let (p_tgt, s_tgt) = decimal_precision_scale(tgt); + if s_src < 0 || s_tgt < 0 { + return false; + } + let src_int = (p_src as i16) - (s_src as i16); + let tgt_int = (p_tgt as i16) - (s_tgt as i16); + tgt_int >= src_int && s_tgt >= s_src +} + +/// Target decimal must be able to represent the full range of the source +/// integer type at the target scale. +fn is_safe_integer_to_decimal(int_type: &DataType, dec_type: &DataType) -> bool { + let required = integer_decimal_digits(int_type); + let (p, s) = decimal_precision_scale(dec_type); + if s < 0 { + return false; + } + (p as i32) - (s as i32) >= required as i32 +} + +/// Minimum number of decimal integer digits needed to store the worst-case +/// value of each integer type without overflow. +fn integer_decimal_digits(int_type: &DataType) -> u8 { + match int_type { + DataType::Int8 | DataType::UInt8 => 3, // 127 / 255 + DataType::Int16 | DataType::UInt16 => 5, // 32767 / 65535 + DataType::Int32 | DataType::UInt32 => 10, // ~2.1e9 / ~4.3e9 + DataType::Int64 => 19, // 9_223_372_036_854_775_807 + DataType::UInt64 => 20, // 18_446_744_073_709_551_615 + _ => unreachable!(), + } +} + +/// Computes a singleton preimage for equality-like predicates over casts whose +/// target value is a string representation of an integer source value. /// /// For example, `CAST(int_col AS Utf8) = '123'` can be rewritten to /// `int_col = 123`, but `CAST(int_col AS Utf8) = '0123'` cannot be rewritten to /// `int_col = 123` because casting `123` back to a string yields `'123'`, not /// `'0123'`. -fn cast_to_string_equality_preimage( +fn exact_preimage_int_to_str_eq_like( source_type: &DataType, target_type: &DataType, op: Operator, @@ -202,7 +383,10 @@ fn cast_to_string_equality_preimage( match (op, lit_value) { ( - Operator::Eq | Operator::NotEq, + Operator::Eq + | Operator::NotEq + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom, ScalarValue::Utf8(Some(_)) | ScalarValue::Utf8View(Some(_)) | ScalarValue::LargeUtf8(Some(_)), @@ -231,12 +415,9 @@ fn cast_to_string_equality_preimage( /// Computes the source-domain preimage interval for timestamp precision /// narrowing casts. /// -/// The preimage is computed entirely in the source timestamp domain and -/// preserves `source_tz`. The target timezone is intentionally *not* copied -/// into the generated bounds: this models the raw timestamp values after -/// truncation and avoids comparing source columns against target-timezone -/// literals. -fn timestamp_precision_narrowing_preimage( +/// Bounds are source-domain literals, so they preserve `source_tz`. The target +/// timezone belongs to the cast result and is not copied into the bounds. +fn timestamp_narrowing_range_preimage( source_type: &DataType, target_type: &DataType, lit_value: &ScalarValue, @@ -2127,4 +2308,216 @@ mod tests { other => panic!("Expected CastPredicatePreimage::Range but got {other:?}"), } } + + // ── Gate safety tests ────────────────────────────────────────────── + + #[test] + fn test_exact_predicate_gate_safe() { + // integer widening + assert_preimage_exact( + &DataType::Int32, + &DataType::Int64, + Operator::Eq, + &ScalarValue::Int64(Some(10)), + ScalarValue::Int32(Some(10)), + ); + // unsigned → wider signed + assert_preimage_exact( + &DataType::UInt32, + &DataType::Int64, + Operator::Eq, + &ScalarValue::Int64(Some(100)), + ScalarValue::UInt32(Some(100)), + ); + // int → decimal (full domain) + assert_preimage_exact( + &DataType::Int32, + &DataType::Decimal128(12, 2), + Operator::Eq, + &ScalarValue::Decimal128(Some(10000), 12, 2), + ScalarValue::Int32(Some(100)), + ); + // uint → decimal (full domain) + assert_preimage_exact( + &DataType::UInt64, + &DataType::Decimal128(20, 0), + Operator::Eq, + &ScalarValue::Decimal128(Some(123), 20, 0), + ScalarValue::UInt64(Some(123)), + ); + // decimal widening + assert_preimage_exact( + &DataType::Decimal128(10, 2), + &DataType::Decimal128(18, 4), + Operator::Eq, + &ScalarValue::Decimal128(Some(1230000), 18, 4), + ScalarValue::Decimal128(Some(12300), 10, 2), + ); + // Date32 ↔ Int32 + assert_preimage_exact( + &DataType::Date32, + &DataType::Int32, + Operator::Eq, + &ScalarValue::Int32(Some(19000)), + ScalarValue::Date32(Some(19000)), + ); + assert_preimage_exact( + &DataType::Int32, + &DataType::Date32, + Operator::Eq, + &ScalarValue::Date32(Some(19000)), + ScalarValue::Int32(Some(19000)), + ); + // Date32 → Date64 is injective; literals must still land on a day + // boundary to be exactly convertible back to the source type. + assert_preimage_exact( + &DataType::Date32, + &DataType::Date64, + Operator::Eq, + &ScalarValue::Date64(Some(1_641_600_000_000)), + ScalarValue::Date32(Some(19000)), + ); + // FixedSizeBinary(n) → Binary + assert_preimage_exact( + &DataType::FixedSizeBinary(4), + &DataType::Binary, + Operator::Eq, + &ScalarValue::Binary(Some(vec![1, 2, 3, 4])), + ScalarValue::FixedSizeBinary(4, Some(vec![1, 2, 3, 4])), + ); + } + + #[test] + fn test_exact_predicate_gate_blocked() { + // numeric narrowing + assert_preimage_none( + &DataType::Int64, + &DataType::Int32, + Operator::Eq, + &ScalarValue::Int32(Some(10)), + ); + // signed → unsigned (even small values) + assert_preimage_none( + &DataType::Int32, + &DataType::UInt32, + Operator::Eq, + &ScalarValue::UInt32(Some(10)), + ); + // unsigned → narrower signed + assert_preimage_none( + &DataType::UInt64, + &DataType::Int64, + Operator::Eq, + &ScalarValue::Int64(Some(10)), + ); + // int → decimal with insufficient precision + assert_preimage_none( + &DataType::Int32, + &DataType::Decimal128(10, 2), + Operator::Eq, + &ScalarValue::Decimal128(Some(10000), 10, 2), + ); + // decimal → int + assert_preimage_none( + &DataType::Decimal128(18, 2), + &DataType::Int64, + Operator::Eq, + &ScalarValue::Int64(Some(123)), + ); + // decimal scale narrowing + assert_preimage_none( + &DataType::Decimal128(18, 2), + &DataType::Decimal128(18, 1), + Operator::Eq, + &ScalarValue::Decimal128(Some(1230), 18, 1), + ); + // negative decimal scale is not accepted for exact predicate rewrites + assert_preimage_none( + &DataType::Decimal128(10, -1), + &DataType::Decimal128(18, 0), + Operator::Eq, + &ScalarValue::Decimal128(Some(120), 18, 0), + ); + assert_preimage_none( + &DataType::Int32, + &DataType::Decimal128(12, -1), + Operator::Eq, + &ScalarValue::Decimal128(Some(120), 12, -1), + ); + // Binary → FixedSizeBinary (not a family-safe source cast) + assert_preimage_none( + &DataType::Binary, + &DataType::FixedSizeBinary(4), + Operator::Eq, + &ScalarValue::FixedSizeBinary(4, Some(vec![1, 2, 3, 4])), + ); + } + + // ── Integer → string distinctness tests ───────────────────────────── + + #[test] + fn test_cast_to_string_is_distinct_from_round_trip() { + // IS DISTINCT FROM with round-trippable string + assert_preimage_exact( + &DataType::Int32, + &DataType::Utf8, + Operator::IsDistinctFrom, + &ScalarValue::Utf8(Some("123".to_string())), + ScalarValue::Int32(Some(123)), + ); + // IS NOT DISTINCT FROM with round-trippable string + assert_preimage_exact( + &DataType::Int32, + &DataType::Utf8, + Operator::IsNotDistinctFrom, + &ScalarValue::Utf8(Some("123".to_string())), + ScalarValue::Int32(Some(123)), + ); + // IS NOT DISTINCT FROM with non-round-trippable string (leading zero) + assert_preimage_none( + &DataType::Int32, + &DataType::Utf8, + Operator::IsNotDistinctFrom, + &ScalarValue::Utf8(Some("0123".to_string())), + ); + // IS DISTINCT FROM with non-round-trippable string + assert_preimage_none( + &DataType::Int32, + &DataType::Utf8, + Operator::IsDistinctFrom, + &ScalarValue::Utf8(Some("0123".to_string())), + ); + } + + // ── Helpers ───────────────────────────────────────────────────────── + + fn assert_preimage_exact( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, + expected_source: ScalarValue, + ) { + let result = + cast_predicate_preimage(source_type, target_type, op, lit_value).unwrap(); + assert_eq!( + result, + Some(CastPredicatePreimage::Exact(expected_source)), + "expected Exact preimage for {source_type:?} → {target_type:?} {op:?} {lit_value:?}" + ); + } + + fn assert_preimage_none( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, + ) { + let result = + cast_predicate_preimage(source_type, target_type, op, lit_value).unwrap(); + assert_eq!( + result, None, + "expected None preimage for {source_type:?} → {target_type:?} {op:?} {lit_value:?}" + ); + } } diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index b03fc44abb45a..bc7db5a18ce1f 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -29,7 +29,7 @@ use datafusion_expr::{ BinaryExpr, Cast, Expr, Operator, TryCast, lit, simplify::SimplifyContext, }; use datafusion_expr_common::casts::{ - CastPredicatePreimage, cast_predicate_exact_literal, cast_predicate_preimage, + CastPredicatePreimage, cast_predicate_preimage, exact_preimage_cast, }; use super::udf_preimage::rewrite_with_preimage; @@ -116,9 +116,12 @@ pub(super) fn supports_cast_predicate_for_inlist( return false; }; + // IN-list rewrites only support singleton exact cast preimages. They do + // not use range preimages (for example timestamp precision narrowing) or + // binary-operator-only special cases such as integer-to-string equality. list.iter().all(|right| match right { Expr::Literal(lit_val, _) => { - cast_predicate_exact_literal(&source_type, target_type, lit_val).is_some() + exact_preimage_cast(&source_type, target_type, lit_val).is_some() } _ => false, }) @@ -140,7 +143,7 @@ pub(super) fn rewrite_cast_predicate_for_inlist( .map(|right| match right { Expr::Literal(lit_value, _) => { let Some(value) = - cast_predicate_exact_literal(&source_type, &target_type, &lit_value) + exact_preimage_cast(&source_type, &target_type, &lit_value) else { return internal_err!( "Can't cast the list expr {:?} to type {}", @@ -195,7 +198,7 @@ mod tests { use datafusion_expr::{binary_expr, cast, col, in_list}; #[test] - fn test_cast_predicate_exact_literal_unwrap() { + fn test_exact_preimage_cast_unwrap() { let schema = expr_test_schema(); let expr = cast(col("c1"), DataType::Int64).gt(lit(10_i64)); @@ -403,6 +406,46 @@ mod tests { assert_eq!(optimize_test(expr, &schema), expected); } + #[test] + fn test_cast_predicate_unsafe_narrowing_kept() { + let schema = expr_test_schema(); + + // CAST(c2 AS Int32) = 5 — should NOT be unwrapped because + // Int64→Int32 is narrowing (blocked by the family gate). + let expr = cast(col("c2"), DataType::Int32).eq(lit(5_i32)); + assert_eq!(optimize_test(expr.clone(), &schema), expr); + + // CAST(c1 AS Int64) = 5 — should still be unwrapped + // (Int32→Int64 is widening, safe). + let expr = cast(col("c1"), DataType::Int64).eq(lit(5_i64)); + let expected = col("c1").eq(lit(5_i32)); + assert_eq!(optimize_test(expr, &schema), expected); + } + + #[test] + fn test_cast_predicate_int_to_string_distinctness() { + let schema = expr_test_schema(); + + // CAST(c1 AS Utf8) IS NOT DISTINCT FROM '123' + let expr = binary_expr( + cast(col("c1"), DataType::Utf8), + Operator::IsNotDistinctFrom, + lit("123"), + ); + let expected = binary_expr(col("c1"), Operator::IsNotDistinctFrom, lit(123_i32)); + assert_eq!(optimize_test(expr, &schema), expected); + + // CAST(c1 AS Utf8) IS DISTINCT FROM '0123' + // Round-trip fails (0123 → 123 → "123" ≠ "0123"), so rewrite + // should NOT happen. + let expr = binary_expr( + cast(col("c1"), DataType::Utf8), + Operator::IsDistinctFrom, + lit("0123"), + ); + assert_eq!(optimize_test(expr.clone(), &schema), expr); + } + fn optimize_test(expr: Expr, schema: &DFSchemaRef) -> Expr { let simplifier = ExprSimplifier::new( SimplifyContext::builder() @@ -418,6 +461,7 @@ mod tests { DFSchema::from_unqualified_fields( vec![ Field::new("c1", DataType::Int32, false), + Field::new("c2", DataType::Int64, false), Field::new("ts_milli", timestamp_millis_type(), false), Field::new("ts_nano", timestamp_nano_type(), false), ] diff --git a/datafusion/physical-expr/src/simplifier/mod.rs b/datafusion/physical-expr/src/simplifier/mod.rs index af87ce7f61d35..dd935effc9c51 100644 --- a/datafusion/physical-expr/src/simplifier/mod.rs +++ b/datafusion/physical-expr/src/simplifier/mod.rs @@ -154,6 +154,8 @@ mod tests { let simplifier = PhysicalExprSimplifier::new(&schema); // Create: cast(c2 as INT32) != INT32(99) + // c2 is Int64 → Int64→Int32 is narrowing → blocked by family gate. + // The cast should stay. let column_expr = col("c2", &schema).unwrap(); let cast_expr = Arc::new(CastExpr::new(column_expr, DataType::Int32, None)); let literal_expr = lit(ScalarValue::Int32(Some(99))); @@ -163,16 +165,14 @@ mod tests { // Apply full simplification (uses TreeNodeRewriter) let optimized = simplifier.simplify(binary_expr).unwrap(); - let optimized_binary = as_binary(&optimized); - - // Should be optimized to: c2 != INT64(99) (c2 is INT64, literal cast to match) - let left_expr = optimized_binary.left(); + // With the family gate, Int64→Int32 is blocked and the cast remains. + let result_binary = as_binary(&optimized); + let result_left = result_binary.left(); assert!( - left_expr.downcast_ref::().is_none() - && left_expr.downcast_ref::().is_none() + result_left.downcast_ref::().is_some() + || result_left.downcast_ref::().is_some(), + "CAST(c2 AS Int32) should remain as Int64→Int32 is narrowing" ); - let right_literal = as_literal(optimized_binary.right()); - assert_eq!(right_literal.value(), &ScalarValue::Int64(Some(99))); } #[test] @@ -198,7 +198,7 @@ mod tests { let or_binary = as_binary(&optimized); - // Verify left side: c1 > INT32(5) + // Verify left side: c1 > INT32(5) (Int32→Int64 widening, unwrapped) let left_binary = as_binary(or_binary.left()); let left_left_expr = left_binary.left(); assert!( @@ -208,15 +208,15 @@ mod tests { let left_literal = as_literal(left_binary.right()); assert_eq!(left_literal.value(), &ScalarValue::Int32(Some(5))); - // Verify right side: c2 <= INT64(10) + // Verify right side: CAST(c2 AS Int32) <= 10 — Int64→Int32 is + // narrowing, blocked by the family gate, so the cast stays. let right_binary = as_binary(or_binary.right()); let right_left_expr = right_binary.left(); assert!( - right_left_expr.downcast_ref::().is_none() - && right_left_expr.downcast_ref::().is_none() + right_left_expr.downcast_ref::().is_some() + || right_left_expr.downcast_ref::().is_some(), + "CAST(c2 AS Int32) should remain as Int64→Int32 is narrowing" ); - let right_literal = as_literal(right_binary.right()); - assert_eq!(right_literal.value(), &ScalarValue::Int64(Some(10))); } #[test] diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index 4b6457422ff46..c408a2b6c756a 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -876,6 +876,8 @@ mod tests { // Create a more complex expression with nested casts // (cast(c1 as INT64) > INT64(10)) AND (cast(c2 as INT32) = INT32(20)) + // c2 is Int64 → the Int64→Int32 cast is narrowing and blocked + // by the family gate, so only the left side should be unwrapped. let c1_expr = col("c1", &schema).unwrap(); let c1_cast = Arc::new(CastExpr::new(c1_expr, DataType::Int64, None)); let c1_literal = lit(10i64); @@ -894,23 +896,22 @@ mod tests { .transform_down(|node| unwrap_cast_in_comparison(node, &schema)) .unwrap(); - // Should be transformed + // Should be transformed (at least the left side) assert!(result.transformed); - // Verify both sides of the AND were optimized + // Verify both sides of the AND let optimized = result.data; let and_binary = optimized.downcast_ref::().unwrap(); - // Left side should be: c1 > INT32(10) + // Left side should be: c1 > INT32(10) (Int32→Int64 widening, unwrapped) let left_binary = and_binary.left().downcast_ref::().unwrap(); assert!(!is_cast_expr(left_binary.left())); let left_literal = left_binary.right().downcast_ref::().unwrap(); assert_eq!(left_literal.value(), &ScalarValue::Int32(Some(10))); - // Right side should be: c2 = INT64(20) (c2 is already INT64, literal cast to match) + // Right side CAST(c2 as INT32) = 20 should stay as-is + // (Int64→Int32 is narrowing, blocked by the family gate). let right_binary = and_binary.right().downcast_ref::().unwrap(); - assert!(!is_cast_expr(right_binary.left())); - let right_literal = right_binary.right().downcast_ref::().unwrap(); - assert_eq!(right_literal.value(), &ScalarValue::Int64(Some(20))); + assert!(is_cast_expr(right_binary.left())); } } From 23948b882225339cf8c2eb3c2c710c9b70a530ef Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 17 Jun 2026 19:14:57 +0800 Subject: [PATCH 07/14] test: helper Signed-off-by: discord9 --- datafusion/expr-common/src/casts.rs | 182 ++++++------------ .../src/simplify_expressions/cast_preimage.rs | 178 +++++++++-------- .../src/simplifier/unwrap_cast.rs | 150 ++++++--------- 3 files changed, 218 insertions(+), 292 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 0aaf827eac1a6..cb2cca64f0b0b 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -1668,95 +1668,39 @@ mod tests { #[test] fn test_cast_predicate_preimage_exact() { - assert_eq!( - cast_predicate_preimage( - &DataType::Int32, - &DataType::Int64, - Operator::Gt, - &ScalarValue::Int64(Some(10)), - ) - .unwrap(), - Some(CastPredicatePreimage::Exact(ScalarValue::Int32(Some(10)))) + assert_preimage_exact( + &DataType::Int32, + &DataType::Int64, + Operator::Gt, + &ScalarValue::Int64(Some(10)), + ScalarValue::Int32(Some(10)), ); - assert_eq!( - cast_predicate_preimage( - &DataType::Int32, - &DataType::Utf8, - Operator::Eq, - &ScalarValue::Utf8(Some("123".to_string())), - ) - .unwrap(), - Some(CastPredicatePreimage::Exact(ScalarValue::Int32(Some(123)))) + assert_preimage_exact( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("123".to_string())), + ScalarValue::Int32(Some(123)), ); - assert_eq!( - cast_predicate_preimage( - &DataType::Int32, - &DataType::Utf8, - Operator::Eq, - &ScalarValue::Utf8(Some("0123".to_string())), - ) - .unwrap(), - None + assert_preimage_none( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("0123".to_string())), ); } #[test] fn test_cast_predicate_preimage_timestamp_narrowing_range() { - let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); - let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); - - assert_eq!( - cast_predicate_preimage( - &ts_ns, - &ts_ms, - Operator::Eq, - &ScalarValue::TimestampMillisecond(Some(1000), None), - ) - .unwrap(), - Some(CastPredicatePreimage::Range( - Interval::try_new( - ScalarValue::TimestampNanosecond(Some(1_000_000_000), None), - ScalarValue::TimestampNanosecond(Some(1_001_000_000), None), - ) - .unwrap() - )) - ); - - assert_eq!( - cast_predicate_preimage( - &ts_ns, - &ts_ms, - Operator::Eq, - &ScalarValue::TimestampMillisecond(Some(0), None), - ) - .unwrap(), - Some(CastPredicatePreimage::Range( - Interval::try_new( - ScalarValue::TimestampNanosecond(Some(-999_999), None), - ScalarValue::TimestampNanosecond(Some(1_000_000), None), - ) - .unwrap() - )) - ); - - assert_eq!( - cast_predicate_preimage( - &ts_ns, - &ts_ms, - Operator::Eq, - &ScalarValue::TimestampMillisecond(Some(-1), None), - ) - .unwrap(), - Some(CastPredicatePreimage::Range( - Interval::try_new( - ScalarValue::TimestampNanosecond(Some(-1_999_999), None), - ScalarValue::TimestampNanosecond(Some(-999_999), None), - ) - .unwrap() - )) - ); + for (lit_ms, lower_ns, upper_ns) in [ + (1000, 1_000_000_000, 1_001_000_000), + (0, -999_999, 1_000_000), + (-1, -1_999_999, -999_999), + ] { + assert_timestamp_narrowing_range(lit_ms, lower_ns, upper_ns); + } } #[test] @@ -1764,28 +1708,19 @@ mod tests { let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); - assert_eq!( - cast_predicate_preimage( - &ts_ms, - &ts_ns, - Operator::Eq, - &ScalarValue::TimestampNanosecond(Some(123_000_000), None), - ) - .unwrap(), - Some(CastPredicatePreimage::Exact( - ScalarValue::TimestampMillisecond(Some(123), None) - )) + assert_preimage_exact( + &ts_ms, + &ts_ns, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_000_000), None), + ScalarValue::TimestampMillisecond(Some(123), None), ); - assert_eq!( - cast_predicate_preimage( - &ts_ms, - &ts_ns, - Operator::Eq, - &ScalarValue::TimestampNanosecond(Some(123_456_789), None), - ) - .unwrap(), - None + assert_preimage_none( + &ts_ms, + &ts_ns, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_456_789), None), ); } @@ -2248,30 +2183,16 @@ mod tests { let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); - // i64::MAX in milliseconds expands beyond i64 range in nanoseconds, + // These millisecond values expand beyond i64 range in nanoseconds, // so the preimage should return None rather than panicking. - assert_eq!( - cast_predicate_preimage( + for value in [i64::MAX, i64::MIN] { + assert_preimage_none( &ts_ns, &ts_ms, Operator::Eq, - &ScalarValue::TimestampMillisecond(Some(i64::MAX), None), - ) - .unwrap(), - None - ); - - // i64::MIN in milliseconds expands beyond i64 range in nanoseconds. - assert_eq!( - cast_predicate_preimage( - &ts_ns, - &ts_ms, - Operator::Eq, - &ScalarValue::TimestampMillisecond(Some(i64::MIN), None), - ) - .unwrap(), - None - ); + &ScalarValue::TimestampMillisecond(Some(value), None), + ); + } } #[test] @@ -2520,4 +2441,25 @@ mod tests { "expected None preimage for {source_type:?} → {target_type:?} {op:?} {lit_value:?}" ); } + + fn assert_timestamp_narrowing_range(lit_ms: i64, lower_ns: i64, upper_ns: i64) { + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + assert_eq!( + cast_predicate_preimage( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(lit_ms), None), + ) + .unwrap(), + Some(CastPredicatePreimage::Range( + Interval::try_new( + ScalarValue::TimestampNanosecond(Some(lower_ns), None), + ScalarValue::TimestampNanosecond(Some(upper_ns), None), + ) + .unwrap() + )) + ); + } } diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index bc7db5a18ce1f..9a5f80596be02 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -285,67 +285,79 @@ mod tests { fn test_cast_preimage_timestamp_precision_narrowing_inequality() { let schema = expr_test_schema(); - let expr = - cast(col("ts_nano"), timestamp_millis_type()).gt(lit_timestamp_millis(1000)); - let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).lt_eq(lit_timestamp_millis(-1)); - let expected = col("ts_nano").lt(lit_timestamp_nano(-999_999)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).lt(lit_timestamp_millis(0)); - let expected = col("ts_nano").lt(lit_timestamp_nano(-999_999)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).lt_eq(lit_timestamp_millis(0)); - let expected = col("ts_nano").lt(lit_timestamp_nano(1_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).gt(lit_timestamp_millis(0)); - let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).gt_eq(lit_timestamp_millis(0)); - let expected = col("ts_nano").gt_eq(lit_timestamp_nano(-999_999)); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = - cast(col("ts_nano"), timestamp_millis_type()).not_eq(lit_timestamp_millis(0)); - let expected = col("ts_nano") - .lt(lit_timestamp_nano(-999_999)) - .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))); - assert_eq!(optimize_test(expr, &schema), expected); + for (op, lit_ms, expected) in vec![ + ( + Operator::Gt, + 1000, + col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)), + ), + ( + Operator::LtEq, + -1, + col("ts_nano").lt(lit_timestamp_nano(-999_999)), + ), + ( + Operator::Lt, + 0, + col("ts_nano").lt(lit_timestamp_nano(-999_999)), + ), + ( + Operator::LtEq, + 0, + col("ts_nano").lt(lit_timestamp_nano(1_000_000)), + ), + ( + Operator::Gt, + 0, + col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000)), + ), + ( + Operator::GtEq, + 0, + col("ts_nano").gt_eq(lit_timestamp_nano(-999_999)), + ), + ( + Operator::NotEq, + 0, + col("ts_nano") + .lt(lit_timestamp_nano(-999_999)) + .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))), + ), + ] { + let expr = binary_expr( + cast(col("ts_nano"), timestamp_millis_type()), + op, + lit_timestamp_millis(lit_ms), + ); + assert_eq!(optimize_test(expr, &schema), expected); + } } #[test] fn test_cast_preimage_timestamp_precision_narrowing_distinctness() { let schema = expr_test_schema(); - let expr = binary_expr( - cast(col("ts_nano"), timestamp_millis_type()), - Operator::IsNotDistinctFrom, - lit_timestamp_millis(0), - ); - let expected = col("ts_nano") - .gt_eq(lit_timestamp_nano(-999_999)) - .and(col("ts_nano").lt(lit_timestamp_nano(1_000_000))); - assert_eq!(optimize_test(expr, &schema), expected); - - let expr = binary_expr( - cast(col("ts_nano"), timestamp_millis_type()), - Operator::IsDistinctFrom, - lit_timestamp_millis(0), - ); - let expected = col("ts_nano") - .lt(lit_timestamp_nano(-999_999)) - .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))); - assert_eq!(optimize_test(expr, &schema), expected); + for (op, expected) in vec![ + ( + Operator::IsNotDistinctFrom, + col("ts_nano") + .gt_eq(lit_timestamp_nano(-999_999)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_000_000))), + ), + ( + Operator::IsDistinctFrom, + col("ts_nano") + .lt(lit_timestamp_nano(-999_999)) + .or(col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000))), + ), + ] { + let expr = binary_expr( + cast(col("ts_nano"), timestamp_millis_type()), + op, + lit_timestamp_millis(0), + ); + assert_eq!(optimize_test(expr, &schema), expected); + } } #[test] @@ -375,35 +387,33 @@ mod tests { fn test_cast_preimage_timestamp_literal_left_range() { let schema = expr_test_schema(); - // lit_timestamp_millis(1000) < cast(col("ts_nano"), timestamp_millis_type()) - // should swap and rewrite to ts_nano >= 1_001_000_000ns - let expr = - lit_timestamp_millis(1000).lt(cast(col("ts_nano"), timestamp_millis_type())); - let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - // lit_timestamp_millis(1000) <= cast(col("ts_nano"), timestamp_millis_type()) - // should swap and rewrite to ts_nano >= 1_000_000_000ns - let expr = lit_timestamp_millis(1000) - .lt_eq(cast(col("ts_nano"), timestamp_millis_type())); - let expected = col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - // lit_timestamp_millis(1000) > cast(col("ts_nano"), timestamp_millis_type()) - // should swap and rewrite to ts_nano < 1_000_000_000ns - let expr = - lit_timestamp_millis(1000).gt(cast(col("ts_nano"), timestamp_millis_type())); - let expected = col("ts_nano").lt(lit_timestamp_nano(1_000_000_000)); - assert_eq!(optimize_test(expr, &schema), expected); - - // lit_timestamp_millis(1000) = cast(col("ts_nano"), timestamp_millis_type()) - // should swap and rewrite to range preimage - let expr = - lit_timestamp_millis(1000).eq(cast(col("ts_nano"), timestamp_millis_type())); - let expected = col("ts_nano") - .gt_eq(lit_timestamp_nano(1_000_000_000)) - .and(col("ts_nano").lt(lit_timestamp_nano(1_001_000_000))); - assert_eq!(optimize_test(expr, &schema), expected); + for (op, expected) in vec![ + ( + Operator::Lt, + col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)), + ), + ( + Operator::LtEq, + col("ts_nano").gt_eq(lit_timestamp_nano(1_000_000_000)), + ), + ( + Operator::Gt, + col("ts_nano").lt(lit_timestamp_nano(1_000_000_000)), + ), + ( + Operator::Eq, + col("ts_nano") + .gt_eq(lit_timestamp_nano(1_000_000_000)) + .and(col("ts_nano").lt(lit_timestamp_nano(1_001_000_000))), + ), + ] { + let expr = binary_expr( + lit_timestamp_millis(1000), + op, + cast(col("ts_nano"), timestamp_millis_type()), + ); + assert_eq!(optimize_test(expr, &schema), expected); + } } #[test] diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index c408a2b6c756a..d2cfcdb02370b 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -237,6 +237,29 @@ mod tests { ]) } + fn timestamp_schema(unit: TimeUnit) -> Schema { + Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(unit, None), + false, + )]) + } + + fn timestamp_cast_comparison( + schema: &Schema, + target_unit: TimeUnit, + op: Operator, + literal: ScalarValue, + ) -> Arc { + let column_expr = col("ts", schema).unwrap(); + let cast_expr = Arc::new(CastExpr::new( + column_expr, + DataType::Timestamp(target_unit, None), + None, + )); + Arc::new(BinaryExpr::new(cast_expr, op, lit(literal))) + } + #[test] fn test_unwrap_cast_in_binary_comparison() { let schema = test_schema(); @@ -636,21 +659,13 @@ mod tests { #[test] fn test_timestamp_precision_narrowing_range_preimage_gt() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )]); - - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Millisecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); - let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::Gt, literal_expr)); + let schema = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, + Operator::Gt, + ScalarValue::TimestampMillisecond(Some(1000), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); @@ -667,21 +682,13 @@ mod tests { #[test] fn test_timestamp_precision_narrowing_range_preimage_eq() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )]); - - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Millisecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(-1), None)); - let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr)); + let schema = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, + Operator::Eq, + ScalarValue::TimestampMillisecond(Some(-1), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); @@ -708,21 +715,13 @@ mod tests { #[test] fn test_timestamp_widening_exactness() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Millisecond, None), - false, - )]); - - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Nanosecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(123_000_000), None)); - let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::GtEq, literal_expr)); + let schema = timestamp_schema(TimeUnit::Millisecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + Operator::GtEq, + ScalarValue::TimestampNanosecond(Some(123_000_000), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); @@ -736,15 +735,12 @@ mod tests { &ScalarValue::TimestampMillisecond(Some(123), None) ); - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Nanosecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(123_456_789), None)); - let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::GtEq, literal_expr)); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + Operator::GtEq, + ScalarValue::TimestampNanosecond(Some(123_456_789), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); assert!(!result.transformed); @@ -752,24 +748,13 @@ mod tests { #[test] fn test_timestamp_precision_narrowing_range_preimage_is_distinct_from() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )]); - - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Millisecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); - let binary_expr = Arc::new(BinaryExpr::new( - cast_expr, + let schema = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, Operator::IsDistinctFrom, - literal_expr, - )); + ScalarValue::TimestampMillisecond(Some(1000), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); @@ -812,24 +797,13 @@ mod tests { #[test] fn test_timestamp_precision_narrowing_range_preimage_is_not_distinct_from() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )]); - - let column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, - DataType::Timestamp(TimeUnit::Millisecond, None), - None, - )); - let literal_expr = lit(ScalarValue::TimestampMillisecond(Some(1000), None)); - let binary_expr = Arc::new(BinaryExpr::new( - cast_expr, + let schema = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, Operator::IsNotDistinctFrom, - literal_expr, - )); + ScalarValue::TimestampMillisecond(Some(1000), None), + ); let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); From c5ed82d12ee139b3c7f468fb1af3992b49d46d72 Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 17 Jun 2026 19:49:50 +0800 Subject: [PATCH 08/14] chore: clippy Signed-off-by: discord9 --- .../optimizer/src/simplify_expressions/cast_preimage.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index 9a5f80596be02..e5c23b945375a 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -285,7 +285,7 @@ mod tests { fn test_cast_preimage_timestamp_precision_narrowing_inequality() { let schema = expr_test_schema(); - for (op, lit_ms, expected) in vec![ + for (op, lit_ms, expected) in [ ( Operator::Gt, 1000, @@ -337,7 +337,7 @@ mod tests { fn test_cast_preimage_timestamp_precision_narrowing_distinctness() { let schema = expr_test_schema(); - for (op, expected) in vec![ + for (op, expected) in [ ( Operator::IsNotDistinctFrom, col("ts_nano") @@ -387,7 +387,7 @@ mod tests { fn test_cast_preimage_timestamp_literal_left_range() { let schema = expr_test_schema(); - for (op, expected) in vec![ + for (op, expected) in [ ( Operator::Lt, col("ts_nano").gt_eq(lit_timestamp_nano(1_001_000_000)), From 7576b32bde49b5d9cf220eada9e0773ca29225fe Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 29 Jul 2026 16:53:28 +0800 Subject: [PATCH 09/14] test: separate cast preimage sqllogictest section --- datafusion/sqllogictest/test_files/simplify_expr.slt | 1 + 1 file changed, 1 insertion(+) diff --git a/datafusion/sqllogictest/test_files/simplify_expr.slt b/datafusion/sqllogictest/test_files/simplify_expr.slt index 76136596aef95..14abb02da018d 100644 --- a/datafusion/sqllogictest/test_files/simplify_expr.slt +++ b/datafusion/sqllogictest/test_files/simplify_expr.slt @@ -499,6 +499,7 @@ select id from date_unwrap where cast(d64 as date) = DATE '1970-01-01' order by statement ok drop table date_unwrap; + # IN-list and array_has simplifications must preserve SQL NULL semantics. query BBB SELECT From f65dd5e536dc711ed0ac3d78b0109be393c4f694 Mon Sep 17 00:00:00 2001 From: discord9 Date: Thu, 30 Jul 2026 02:17:43 +0800 Subject: [PATCH 10/14] fix: rewrite ordered timestamp widening predicates --- datafusion/expr-common/src/casts.rs | 222 +++++++++++++++++- .../src/simplify_expressions/cast_preimage.rs | 75 +++++- .../src/simplifier/unwrap_cast.rs | 132 +++++++++-- 3 files changed, 397 insertions(+), 32 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index cb2cca64f0b0b..a158351ba5b33 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -40,8 +40,11 @@ use datafusion_common::{Result, ScalarValue}; /// Source-domain preimage of `CAST(source_expr AS target_type) OP literal`. #[derive(Debug, Clone, PartialEq, Eq)] pub enum CastPredicatePreimage { - /// A singleton preimage represented by a literal in the source type. This - /// can keep the original comparison operator. + /// An exact same-operator source bound or literal. + /// + /// For equality-like predicates this is a singleton literal. For ordered + /// timestamp widening predicates it is a source-unit bound, not a + /// singleton equality preimage. Exact(ScalarValue), /// A half-open source-domain interval `[lower, upper)`. The caller must /// map the comparison operator to range predicates. @@ -94,8 +97,9 @@ pub fn try_cast_literal_to_type( /// Computes a source-domain preimage for `CAST(source AS target_type) OP literal`. /// /// This is the shared semantic core for logical and physical cast-predicate -/// rewrites. It returns a singleton [`CastPredicatePreimage::Exact`] for casts -/// where moving the cast to the literal preserves comparison semantics, and a +/// rewrites. It returns an exact same-operator source bound or literal as +/// [`CastPredicatePreimage::Exact`] where moving the cast preserves comparison +/// semantics, and a /// [`CastPredicatePreimage::Range`] for many-to-one casts with known preimages /// such as timestamp precision narrowing. pub fn cast_predicate_preimage( @@ -118,8 +122,14 @@ pub fn cast_predicate_preimage( return Ok(Some(CastPredicatePreimage::Exact(value))); } - Ok(exact_preimage_cast(source_type, target_type, lit_value) - .map(CastPredicatePreimage::Exact)) + if let Some(value) = exact_preimage_cast(source_type, target_type, lit_value) { + return Ok(Some(CastPredicatePreimage::Exact(value))); + } + + Ok( + timestamp_widening_ordered_preimage(source_type, target_type, op, lit_value) + .map(CastPredicatePreimage::Exact), + ) } fn maybe_range_preimage( @@ -460,6 +470,63 @@ fn timestamp_narrowing_range_preimage( .map(Some) } +/// Computes a same-operator source bound for an ordered timestamp precision +/// widening cast with a non-aligned target literal. +/// +/// This deliberately follows the accepted timestamp-widening policy for this +/// PR: at extreme source values where widening overflows, `CAST` can error and +/// `TRY_CAST` can return `NULL`, while the rewritten source-unit comparison +/// returns a Boolean. It is therefore not a global full-domain equivalence +/// claim for those overflow cases; a guarded-preimage redesign is out of scope. +fn timestamp_widening_ordered_preimage( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, +) -> Option { + let ( + DataType::Timestamp(source_unit, source_tz), + DataType::Timestamp(target_unit, target_tz), + ) = (source_type, target_type) + else { + return None; + }; + + if source_tz != target_tz + || lit_value.is_null() + || lit_value.data_type() != *target_type + || !matches!( + op, + Operator::Lt | Operator::LtEq | Operator::Gt | Operator::GtEq + ) + { + return None; + } + + let source_scale = i128::from(timestamp_unit_scale(source_unit)); + let target_scale = i128::from(timestamp_unit_scale(target_unit)); + if target_scale <= source_scale { + return None; + } + + let target_value = i128::from(timestamp_literal_value(lit_value, target_unit)?); + let quotient = target_scale / source_scale; + let floor = target_value.div_euclid(quotient); + let remainder = target_value.rem_euclid(quotient); + if remainder == 0 { + return None; + } + let ceil = floor + 1; + let bound = match op { + Operator::GtEq | Operator::Lt => ceil, + Operator::Gt | Operator::LtEq => floor, + _ => return None, + }; + + let bound = i64::try_from(bound).ok()?; + Some(timestamp_scalar(source_unit, source_tz.clone(), bound)) +} + /// Returns the half-open source-domain bucket `[lower, upper)` that truncates /// toward zero to `value` when divided by `bucket_width`. /// @@ -1704,7 +1771,7 @@ mod tests { } #[test] - fn test_cast_predicate_preimage_timestamp_widening_exact_only() { + fn test_cast_predicate_preimage_timestamp_widening_ordered() { let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); @@ -1716,6 +1783,17 @@ mod tests { ScalarValue::TimestampMillisecond(Some(123), None), ); + for (source_unit, target_unit) in [ + (TimeUnit::Second, TimeUnit::Millisecond), + (TimeUnit::Second, TimeUnit::Microsecond), + (TimeUnit::Second, TimeUnit::Nanosecond), + (TimeUnit::Millisecond, TimeUnit::Microsecond), + (TimeUnit::Millisecond, TimeUnit::Nanosecond), + (TimeUnit::Microsecond, TimeUnit::Nanosecond), + ] { + assert_timestamp_widening_ordered(source_unit, target_unit); + } + assert_preimage_none( &ts_ms, &ts_ns, @@ -1724,6 +1802,59 @@ mod tests { ); } + #[test] + fn test_timestamp_widening_ordered_preimage_rejections_and_timezone() { + let source_type = DataType::Timestamp(TimeUnit::Millisecond, Some("UTC".into())); + let target_type = DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into())); + let literal = + ScalarValue::TimestampNanosecond(Some(123_456_789), Some("UTC".into())); + + assert_preimage_exact( + &source_type, + &target_type, + Operator::GtEq, + &literal, + ScalarValue::TimestampMillisecond(Some(124), Some("UTC".into())), + ); + + for op in [ + Operator::Eq, + Operator::NotEq, + Operator::IsDistinctFrom, + Operator::IsNotDistinctFrom, + ] { + assert_preimage_none(&source_type, &target_type, op, &literal); + } + + assert!( + timestamp_widening_ordered_preimage( + &source_type, + &target_type, + Operator::GtEq, + &ScalarValue::TimestampNanosecond(None, Some("UTC".into())), + ) + .is_none() + ); + assert!( + timestamp_widening_ordered_preimage( + &source_type, + &target_type, + Operator::GtEq, + &ScalarValue::TimestampMicrosecond(Some(123_456), Some("UTC".into())), + ) + .is_none() + ); + assert!( + timestamp_widening_ordered_preimage( + &source_type, + &DataType::Timestamp(TimeUnit::Nanosecond, None), + Operator::GtEq, + &ScalarValue::TimestampNanosecond(Some(123_456_789), None), + ) + .is_none() + ); + } + #[test] fn test_cast_predicate_preimage_timestamp_null_literal_unsupported() { let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); @@ -2442,6 +2573,83 @@ mod tests { ); } + fn assert_timestamp_widening_ordered(source_unit: TimeUnit, target_unit: TimeUnit) { + let source_type = DataType::Timestamp(source_unit, None); + let target_type = DataType::Timestamp(target_unit, None); + let quotient = i128::from(timestamp_unit_scale(&target_unit)) + / i128::from(timestamp_unit_scale(&source_unit)); + + for (target_value, floor, ceil) in [ + (i64::try_from(quotient + 1).unwrap(), 1, 2), + (i64::try_from(1 - quotient).unwrap(), -1, 0), + ] { + for (op, expected) in [ + (Operator::GtEq, ceil), + (Operator::Gt, floor), + (Operator::Lt, ceil), + (Operator::LtEq, floor), + ] { + assert_preimage_exact( + &source_type, + &target_type, + op, + ×tamp_scalar(&target_unit, None, target_value), + timestamp_scalar(&source_unit, None, expected), + ); + } + } + + // Aligned literals continue through the exact-preimage path. + for op in [Operator::GtEq, Operator::Gt, Operator::Lt, Operator::LtEq] { + assert_preimage_exact( + &source_type, + &target_type, + op, + ×tamp_scalar( + &target_unit, + None, + i64::try_from(123 * quotient).unwrap(), + ), + timestamp_scalar(&source_unit, None, 123), + ); + } + + // The accepted widening policy applies the Euclidean source bound to + // boundary literals as well, despite CAST/TRY_CAST overflow caveats. + for target_value in [i64::MIN, i64::MAX] { + let target_value = i128::from(target_value); + let floor = target_value.div_euclid(quotient); + let ceil = floor + + if target_value.rem_euclid(quotient) != 0 { + 1 + } else { + 0 + }; + for (op, expected) in [ + (Operator::GtEq, ceil), + (Operator::Gt, floor), + (Operator::Lt, ceil), + (Operator::LtEq, floor), + ] { + assert_preimage_exact( + &source_type, + &target_type, + op, + ×tamp_scalar( + &target_unit, + None, + i64::try_from(target_value).unwrap(), + ), + timestamp_scalar( + &source_unit, + None, + i64::try_from(expected).unwrap(), + ), + ); + } + } + } + fn assert_timestamp_narrowing_range(lit_ms: i64, lower_ns: i64, upper_ns: i64) { let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index e5c23b945375a..b1832c833ea27 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -195,7 +195,7 @@ mod tests { use arrow::datatypes::{Field, TimeUnit}; use datafusion_common::{DFSchema, DFSchemaRef, ScalarValue}; use datafusion_expr::simplify::SimplifyContext; - use datafusion_expr::{binary_expr, cast, col, in_list}; + use datafusion_expr::{binary_expr, cast, col, in_list, try_cast}; #[test] fn test_exact_preimage_cast_unwrap() { @@ -361,7 +361,7 @@ mod tests { } #[test] - fn test_cast_preimage_timestamp_widening_requires_exact_literal() { + fn test_cast_preimage_timestamp_widening_ordered() { let schema = expr_test_schema(); let expr = cast(col("ts_milli"), timestamp_nano_type()) @@ -369,20 +369,83 @@ mod tests { let expected = col("ts_milli").eq(lit_timestamp_millis(123)); assert_eq!(optimize_test(expr, &schema), expected); - let expr = cast(col("ts_milli"), timestamp_nano_type()) - .eq(lit_timestamp_nano(123_456_789)); - assert_eq!(optimize_test(expr.clone(), &schema), expr); + for (literal, floor, ceil) in + [(123_456_789, 123, 124), (-123_456_789, -124, -123)] + { + for (op, bound) in [ + (Operator::GtEq, ceil), + (Operator::Gt, floor), + (Operator::Lt, ceil), + (Operator::LtEq, floor), + ] { + let expr = binary_expr( + cast(col("ts_milli"), timestamp_nano_type()), + op, + lit_timestamp_nano(literal), + ); + let expected = + binary_expr(col("ts_milli"), op, lit_timestamp_millis(bound)); + assert_eq!(optimize_test(expr, &schema), expected); + } + } let expr = cast(col("ts_milli"), timestamp_nano_type()) .gt_eq(lit_timestamp_nano(123_000_000)); let expected = col("ts_milli").gt_eq(lit_timestamp_millis(123)); assert_eq!(optimize_test(expr, &schema), expected); - let expr = cast(col("ts_milli"), timestamp_nano_type()) + let expr = try_cast(col("ts_milli"), timestamp_nano_type()) .gt_eq(lit_timestamp_nano(123_456_789)); + let expected = col("ts_milli").gt_eq(lit_timestamp_millis(124)); + assert_eq!(optimize_test(expr, &schema), expected); + + let expr = cast(col("ts_milli"), timestamp_nano_type()) + .eq(lit_timestamp_nano(123_456_789)); + assert_eq!(optimize_test(expr.clone(), &schema), expr); + + let expr = in_list( + cast(col("ts_milli"), timestamp_nano_type()), + vec![ + lit_timestamp_nano(123_456_789), + lit_timestamp_nano(987_654_321), + ], + false, + ); assert_eq!(optimize_test(expr.clone(), &schema), expr); } + #[test] + fn test_cast_preimage_timestamp_widening_literal_left_ordered() { + let schema = expr_test_schema(); + + for (op, expected_op, expected_value) in [ + (Operator::Lt, Operator::Gt, 123), + (Operator::LtEq, Operator::GtEq, 124), + (Operator::Gt, Operator::Lt, 124), + (Operator::GtEq, Operator::LtEq, 123), + ] { + let expr = binary_expr( + lit_timestamp_nano(123_456_789), + op, + cast(col("ts_milli"), timestamp_nano_type()), + ); + let expected = binary_expr( + col("ts_milli"), + expected_op, + lit_timestamp_millis(expected_value), + ); + assert_eq!(optimize_test(expr, &schema), expected); + } + + let expr = binary_expr( + lit_timestamp_nano(-123_456_789), + Operator::GtEq, + try_cast(col("ts_milli"), timestamp_nano_type()), + ); + let expected = col("ts_milli").lt_eq(lit_timestamp_millis(-124)); + assert_eq!(optimize_test(expr, &schema), expected); + } + #[test] fn test_cast_preimage_timestamp_literal_left_range() { let schema = expr_test_schema(); diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index d2cfcdb02370b..776eec2d110f6 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -260,6 +260,24 @@ mod tests { Arc::new(BinaryExpr::new(cast_expr, op, lit(literal))) } + fn assert_timestamp_widening_rewrite( + schema: &Schema, + binary_expr: Arc, + expected_op: Operator, + expected_value: i64, + ) { + let result = unwrap_cast_in_comparison(binary_expr, schema).unwrap(); + assert!(result.transformed); + let optimized_binary = result.data.downcast_ref::().unwrap(); + assert_eq!(*optimized_binary.op(), expected_op); + assert!(!is_cast_expr(optimized_binary.left())); + let right_literal = optimized_binary.right().downcast_ref::().unwrap(); + assert_eq!( + right_literal.value(), + &ScalarValue::TimestampMillisecond(Some(expected_value), None) + ); + } + #[test] fn test_unwrap_cast_in_binary_comparison() { let schema = test_schema(); @@ -714,38 +732,114 @@ mod tests { } #[test] - fn test_timestamp_widening_exactness() { + fn test_timestamp_widening_ordered() { let schema = timestamp_schema(TimeUnit::Millisecond); - let binary_expr = timestamp_cast_comparison( + assert_timestamp_widening_rewrite( &schema, - TimeUnit::Nanosecond, + timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + Operator::GtEq, + ScalarValue::TimestampNanosecond(Some(123_000_000), None), + ), Operator::GtEq, - ScalarValue::TimestampNanosecond(Some(123_000_000), None), + 123, ); - let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); - - assert!(result.transformed); - let optimized_binary = result.data.downcast_ref::().unwrap(); - assert_eq!(*optimized_binary.op(), Operator::GtEq); - assert!(!is_cast_expr(optimized_binary.left())); - let right_literal = optimized_binary.right().downcast_ref::().unwrap(); - assert_eq!( - right_literal.value(), - &ScalarValue::TimestampMillisecond(Some(123), None) - ); + for (literal, floor, ceil) in + [(123_456_789, 123, 124), (-123_456_789, -124, -123)] + { + for (op, bound) in [ + (Operator::GtEq, ceil), + (Operator::Gt, floor), + (Operator::Lt, ceil), + (Operator::LtEq, floor), + ] { + assert_timestamp_widening_rewrite( + &schema, + timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + op, + ScalarValue::TimestampNanosecond(Some(literal), None), + ), + op, + bound, + ); + } + } - let binary_expr = timestamp_cast_comparison( + let try_cast_expr = Arc::new(TryCastExpr::new( + col("ts", &schema).unwrap(), + DataType::Timestamp(TimeUnit::Nanosecond, None), + )); + assert_timestamp_widening_rewrite( &schema, - TimeUnit::Nanosecond, + Arc::new(BinaryExpr::new( + try_cast_expr, + Operator::GtEq, + lit(ScalarValue::TimestampNanosecond(Some(123_456_789), None)), + )), Operator::GtEq, - ScalarValue::TimestampNanosecond(Some(123_456_789), None), + 124, ); - let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + let result = unwrap_cast_in_comparison( + timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + Operator::Eq, + ScalarValue::TimestampNanosecond(Some(123_456_789), None), + ), + &schema, + ) + .unwrap(); assert!(!result.transformed); } + #[test] + fn test_timestamp_widening_literal_left_ordered() { + let schema = timestamp_schema(TimeUnit::Millisecond); + for (op, expected_op, expected_value) in [ + (Operator::Lt, Operator::Gt, 123), + (Operator::LtEq, Operator::GtEq, 124), + (Operator::Gt, Operator::Lt, 124), + (Operator::GtEq, Operator::LtEq, 123), + ] { + let cast_expr = Arc::new(CastExpr::new( + col("ts", &schema).unwrap(), + DataType::Timestamp(TimeUnit::Nanosecond, None), + None, + )); + let binary_expr = Arc::new(BinaryExpr::new( + lit(ScalarValue::TimestampNanosecond(Some(123_456_789), None)), + op, + cast_expr, + )); + assert_timestamp_widening_rewrite( + &schema, + binary_expr, + expected_op, + expected_value, + ); + } + + let try_cast_expr = Arc::new(TryCastExpr::new( + col("ts", &schema).unwrap(), + DataType::Timestamp(TimeUnit::Nanosecond, None), + )); + assert_timestamp_widening_rewrite( + &schema, + Arc::new(BinaryExpr::new( + lit(ScalarValue::TimestampNanosecond(Some(-123_456_789), None)), + Operator::GtEq, + try_cast_expr, + )), + Operator::LtEq, + -124, + ); + } + #[test] fn test_timestamp_precision_narrowing_range_preimage_is_distinct_from() { let schema = timestamp_schema(TimeUnit::Nanosecond); From a7e0c47b4af5a8dea69cc2f58ce963f264d3b798 Mon Sep 17 00:00:00 2001 From: discord9 Date: Thu, 30 Jul 2026 02:20:59 +0800 Subject: [PATCH 11/14] test: cover ordered timestamp widening plan --- datafusion/sqllogictest/test_files/simplify_expr.slt | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/datafusion/sqllogictest/test_files/simplify_expr.slt b/datafusion/sqllogictest/test_files/simplify_expr.slt index 14abb02da018d..6a12089d61682 100644 --- a/datafusion/sqllogictest/test_files/simplify_expr.slt +++ b/datafusion/sqllogictest/test_files/simplify_expr.slt @@ -564,10 +564,10 @@ EXPLAIN SELECT * FROM cast_preimage_ts WHERE CAST(ts_milli AS TIMESTAMP(9)) >= arrow_cast(123456789, 'Timestamp(ns)'); ---- logical_plan -01)Filter: CAST(cast_preimage_ts.ts_milli AS Timestamp(ns)) >= TimestampNanosecond(123456789, None) +01)Filter: cast_preimage_ts.ts_milli >= TimestampMillisecond(124, None) 02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] physical_plan -01)FilterExec: CAST(ts_milli@1 AS Timestamp(ns)) >= 123456789 +01)FilterExec: ts_milli@1 >= 124 02)--DataSourceExec: partitions=1, partition_sizes=[1] query TT From 08357093f05515a22fb1538ced4aa078551b01a6 Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Fri, 28 Aug 2026 17:51:57 +0800 Subject: [PATCH 12/14] chore: address clippy lints after rebase Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- datafusion/expr-common/src/casts.rs | 35 +++++++++---------- .../src/simplify_expressions/cast_preimage.rs | 8 ++--- 2 files changed, 20 insertions(+), 23 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index a158351ba5b33..948957c5758fd 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -546,19 +546,21 @@ fn timestamp_widening_ordered_preimage( /// intermediate bound cannot be represented as `i128`. fn trunc_toward_zero_bucket(value: i64, bucket_width: i128) -> Option<(i128, i128)> { let value = value as i128; - if value > 0 { - let lower = value.checked_mul(bucket_width)?; - let upper = value.checked_add(1)?.checked_mul(bucket_width)?; - Some((lower, upper)) - } else if value == 0 { - Some((1_i128.checked_sub(bucket_width)?, bucket_width)) - } else { - let lower = value - .checked_sub(1)? - .checked_mul(bucket_width)? - .checked_add(1)?; - let upper = value.checked_mul(bucket_width)?.checked_add(1)?; - Some((lower, upper)) + match value.cmp(&0) { + Ordering::Greater => { + let lower = value.checked_mul(bucket_width)?; + let upper = value.checked_add(1)?.checked_mul(bucket_width)?; + Some((lower, upper)) + } + Ordering::Equal => Some((1_i128.checked_sub(bucket_width)?, bucket_width)), + Ordering::Less => { + let lower = value + .checked_sub(1)? + .checked_mul(bucket_width)? + .checked_add(1)?; + let upper = value.checked_mul(bucket_width)?.checked_add(1)?; + Some((lower, upper)) + } } } @@ -2619,12 +2621,7 @@ mod tests { for target_value in [i64::MIN, i64::MAX] { let target_value = i128::from(target_value); let floor = target_value.div_euclid(quotient); - let ceil = floor - + if target_value.rem_euclid(quotient) != 0 { - 1 - } else { - 0 - }; + let ceil = floor + i128::from(target_value.rem_euclid(quotient) != 0); for (op, expected) in [ (Operator::GtEq, ceil), (Operator::Gt, floor), diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index b1832c833ea27..0ca5597a69fab 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -169,16 +169,16 @@ pub(super) fn rewrite_cast_predicate_for_inlist( fn cast_input_and_type(cast_expr: Expr) -> Option<(Box, DataType)> { match cast_expr { - Expr::TryCast(TryCast { expr, field, .. }) - | Expr::Cast(Cast { expr, field, .. }) => Some((expr, field.data_type().clone())), + Expr::TryCast(TryCast { expr, field }) | Expr::Cast(Cast { expr, field }) => { + Some((expr, field.data_type().clone())) + } _ => None, } } fn cast_input_and_type_ref(cast_expr: &Expr) -> Option<(&Expr, &DataType)> { match cast_expr { - Expr::TryCast(TryCast { expr, field, .. }) - | Expr::Cast(Cast { expr, field, .. }) => { + Expr::TryCast(TryCast { expr, field }) | Expr::Cast(Cast { expr, field }) => { Some((expr.as_ref(), field.data_type())) } _ => None, From 5fbc43ed6f95acc862f538e53206cabfc3a76aec Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Wed, 9 Sep 2026 17:31:43 +0800 Subject: [PATCH 13/14] fix: preserve timezone and volatile cast predicate semantics Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- .../core/tests/expr_api/simplification.rs | 279 +++++++++++++++++- datafusion/expr-common/src/casts.rs | 35 ++- .../src/simplify_expressions/cast_preimage.rs | 22 +- .../src/simplifier/unwrap_cast.rs | 13 + 4 files changed, 332 insertions(+), 17 deletions(-) diff --git a/datafusion/core/tests/expr_api/simplification.rs b/datafusion/core/tests/expr_api/simplification.rs index e9a975239a481..78ecf5365d6b5 100644 --- a/datafusion/core/tests/expr_api/simplification.rs +++ b/datafusion/core/tests/expr_api/simplification.rs @@ -20,25 +20,32 @@ use insta::assert_snapshot; use arrow::array::types::IntervalDayTime; -use arrow::array::{ArrayRef, Int32Array}; +use arrow::array::{ + ArrayRef, BooleanArray, Int32Array, RecordBatch, TimestampNanosecondArray, +}; use arrow::datatypes::{DataType, Field, Schema}; use chrono::{DateTime, TimeZone, Utc}; use datafusion::{error::Result, prelude::*}; use datafusion_common::ScalarValue; use datafusion_common::cast::as_int32_array; use datafusion_common::{DFSchemaRef, ToDFSchema}; -use datafusion_expr::expr::ScalarFunction; +use datafusion_expr::execution_props::ExecutionProps; +use datafusion_expr::expr::{ScalarFunction, TryCast}; use datafusion_expr::logical_plan::builder::table_scan_with_filters; +use datafusion_expr::physical_planning_context::PhysicalPlanningContext; use datafusion_expr::simplify::SimplifyContext; use datafusion_expr::{ - Cast, ColumnarValue, ExprSchemable, LogicalPlan, LogicalPlanBuilder, Projection, - ScalarUDF, Volatility, table_scan, + Cast, ColumnarValue, ExprSchemable, LogicalPlan, LogicalPlanBuilder, Operator, + Projection, ScalarUDF, Volatility, table_scan, }; use datafusion_functions::math; use datafusion_optimizer::optimizer::Optimizer; use datafusion_optimizer::simplify_expressions::{ExprSimplifier, SimplifyExpressions}; use datafusion_optimizer::{OptimizerContext, OptimizerRule}; +use datafusion_physical_expr::PhysicalExprSimplifier; +use datafusion_physical_expr::planner::create_physical_expr; use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; /// A schema like: /// @@ -152,6 +159,270 @@ fn to_timestamp_expr(arg: impl Into) -> Expr { to_timestamp(vec![lit(arg.into())]) } +fn cast_or_try_cast(expr: Expr, data_type: DataType, try_cast: bool) -> Expr { + if try_cast { + Expr::TryCast(TryCast::new(Box::new(expr), data_type)) + } else { + Expr::Cast(Cast::new(Box::new(expr), data_type)) + } +} + +fn simplify_logical_expr(expr: Expr, schema: DFSchemaRef) -> Expr { + ExprSimplifier::new(SimplifyContext::builder().with_schema(schema).build()) + .simplify(expr) + .unwrap() +} + +fn evaluate_physical_boolean_expr( + physical_expr: Arc, + batch: &RecordBatch, +) -> Vec> { + let value = physical_expr.evaluate(batch).unwrap(); + let array = value.into_array(batch.num_rows()).unwrap(); + array + .as_any() + .downcast_ref::() + .unwrap() + .iter() + .collect() +} + +fn evaluate_boolean_expr(expr: &Expr, batch: &RecordBatch) -> Vec> { + let schema = batch.schema().to_dfschema_ref().unwrap(); + let physical_expr = create_physical_expr( + expr, + schema.as_ref(), + &ExecutionProps::new(), + &PhysicalPlanningContext::default(), + ) + .unwrap(); + evaluate_physical_boolean_expr(physical_expr, batch) +} + +fn evaluate_physical_simplified_boolean_expr( + expr: &Expr, + batch: &RecordBatch, +) -> Vec> { + let schema = batch.schema().to_dfschema_ref().unwrap(); + // Lower the unsimplified logical expression so this specifically tests the + // physical simplifier rather than a full SQL/planning pipeline. + let physical_expr = create_physical_expr( + expr, + schema.as_ref(), + &ExecutionProps::new(), + &PhysicalPlanningContext::default(), + ) + .unwrap(); + let simplified = PhysicalExprSimplifier::new(batch.schema().as_ref()) + .simplify(physical_expr) + .unwrap(); + evaluate_physical_boolean_expr(simplified, batch) +} + +#[test] +fn timestamp_timezone_cast_preimage_preserves_results() { + let source_type = DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, None); + let target_type = DataType::Timestamp( + arrow::datatypes::TimeUnit::Millisecond, + Some("+01:00".into()), + ); + let batch = RecordBatch::try_from_iter(vec![( + "ts", + Arc::new(TimestampNanosecondArray::from(vec![ + Some(-1_000_000), + Some(-999_999), + Some(0), + Some(999_999), + Some(1_000_000), + None, + ])) as ArrayRef, + )]) + .unwrap(); + assert_eq!(batch.schema().field(0).data_type(), &source_type); + let expected = vec![ + Some(false), + Some(false), + Some(false), + Some(false), + Some(false), + None, + ]; + + for try_cast in [false, true] { + let make_expr = || { + cast_or_try_cast(col("ts"), target_type.clone(), try_cast).eq(lit( + ScalarValue::TimestampMillisecond(Some(0), Some("+01:00".into())), + )) + }; + let original = make_expr(); + let logical = + simplify_logical_expr(make_expr(), batch.schema().to_dfschema_ref().unwrap()); + let logical_rows = evaluate_boolean_expr(&logical, &batch); + let physical_rows = + evaluate_physical_simplified_boolean_expr(&make_expr(), &batch); + + // Naive timestamps cannot use the timezone-aware target's bucket preimage. + assert_eq!(evaluate_boolean_expr(&original, &batch), expected); + assert_eq!(logical_rows, expected); + assert_eq!(physical_rows, expected); + } +} + +#[test] +fn timestamp_narrowing_range_controls_still_rewrite() { + for timezone in [None, Some("+01:00".into())] { + let source_type = + DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, timezone.clone()); + let target_type = DataType::Timestamp( + arrow::datatypes::TimeUnit::Millisecond, + timezone.clone(), + ); + let batch = RecordBatch::try_from_iter(vec![( + "ts", + Arc::new( + TimestampNanosecondArray::from(vec![ + Some(-1_000_000), + Some(-999_999), + Some(0), + Some(999_999), + Some(1_000_000), + None, + ]) + .with_timezone_opt(timezone.clone()), + ) as ArrayRef, + )]) + .unwrap(); + assert_eq!(batch.schema().field(0).data_type(), &source_type); + let expected = vec![ + Some(false), + Some(true), + Some(true), + Some(true), + Some(false), + None, + ]; + + for try_cast in [false, true] { + let make_expr = || { + cast_or_try_cast(col("ts"), target_type.clone(), try_cast).eq(lit( + ScalarValue::TimestampMillisecond(Some(0), timezone.clone()), + )) + }; + let original = make_expr(); + let logical = simplify_logical_expr( + make_expr(), + batch.schema().to_dfschema_ref().unwrap(), + ); + let physical_rows = + evaluate_physical_simplified_boolean_expr(&make_expr(), &batch); + assert_ne!(logical, original); + assert_eq!(evaluate_boolean_expr(&original, &batch), expected); + assert_eq!(evaluate_boolean_expr(&logical, &batch), expected); + assert_eq!(physical_rows, expected); + } + } +} + +fn alternating_timestamp_udf(counter: Arc) -> Arc { + Arc::new(create_udf( + "alternating_timestamp", + vec![], + DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, None), + Volatility::Volatile, + Arc::new(move |_args: &[ColumnarValue]| { + let value = if counter.fetch_add(1, Ordering::SeqCst).is_multiple_of(2) { + 1_000_000 + } else { + -1_000_000 + }; + Ok(ColumnarValue::Scalar(ScalarValue::TimestampNanosecond( + Some(value), + None, + ))) + }), + )) +} + +#[test] +fn volatile_cast_preimage_does_not_duplicate_evaluation() { + let counter = Arc::new(AtomicUsize::new(0)); + let batch = RecordBatch::try_from_iter(vec![( + "one_row", + Arc::new(Int32Array::from(vec![1])) as ArrayRef, + )]) + .unwrap(); + let target_type = DataType::Timestamp(arrow::datatypes::TimeUnit::Millisecond, None); + + for try_cast in [false, true] { + for op in [ + Operator::Eq, + Operator::NotEq, + Operator::IsDistinctFrom, + Operator::IsNotDistinctFrom, + ] { + let make_expr = || { + binary_expr( + cast_or_try_cast( + Expr::ScalarFunction(ScalarFunction::new_udf( + alternating_timestamp_udf(Arc::clone(&counter)), + vec![], + )), + target_type.clone(), + try_cast, + ), + op, + lit(ScalarValue::TimestampMillisecond(Some(0), None)), + ) + }; + counter.store(0, Ordering::SeqCst); + let original = make_expr(); + let expected = evaluate_boolean_expr(&original, &batch); + assert_eq!(counter.load(Ordering::SeqCst), 1); + + counter.store(0, Ordering::SeqCst); + let logical = simplify_logical_expr( + make_expr(), + batch.schema().to_dfschema_ref().unwrap(), + ); + assert_eq!(evaluate_boolean_expr(&logical, &batch), expected); + assert_eq!(counter.load(Ordering::SeqCst), 1, "logical {op:?}"); + + counter.store(0, Ordering::SeqCst); + let physical = + evaluate_physical_simplified_boolean_expr(&make_expr(), &batch); + assert_eq!(physical, expected); + assert_eq!(counter.load(Ordering::SeqCst), 1, "physical {op:?}"); + } + + // Ordered range preimages still have one input reference and remain enabled. + let make_expr = || { + cast_or_try_cast( + Expr::ScalarFunction(ScalarFunction::new_udf( + alternating_timestamp_udf(Arc::clone(&counter)), + vec![], + )), + target_type.clone(), + try_cast, + ) + .lt(lit(ScalarValue::TimestampMillisecond(Some(0), None))) + }; + counter.store(0, Ordering::SeqCst); + let original = make_expr(); + let expected = evaluate_boolean_expr(&original, &batch); + assert_eq!(counter.load(Ordering::SeqCst), 1); + counter.store(0, Ordering::SeqCst); + let logical = + simplify_logical_expr(make_expr(), batch.schema().to_dfschema_ref().unwrap()); + assert_ne!(logical, original); + assert_eq!(evaluate_boolean_expr(&logical, &batch), expected); + assert_eq!(counter.load(Ordering::SeqCst), 1); + counter.store(0, Ordering::SeqCst); + let physical = evaluate_physical_simplified_boolean_expr(&make_expr(), &batch); + assert_eq!(physical, expected); + assert_eq!(counter.load(Ordering::SeqCst), 1); + } +} + #[test] fn basic() { let context = SimplifyContext::builder() diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 948957c5758fd..a8e8d1dcbcf06 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -160,11 +160,12 @@ pub fn exact_preimage_cast( return None; } - // Apply a family-level safety gate: the source→target cast must be + // Apply a family-level safety gate: the source→target cast must normally be // value-preserving over the full source domain. Timestamp precision - // narrowing is handled as a Range preimage, not an Exact preimage; - // timestamp widening / equal-unit casts still require the literal - // round-trip check below. + // narrowing is handled as a Range preimage, not an Exact preimage. + // Timestamp widening / equal-unit casts still require the literal + // round-trip check below; widening retains its pre-existing overflow + // limitation at extreme source values. if !is_exact_cast_safe(source_type, target_type) { return None; } @@ -203,14 +204,17 @@ fn is_timestamp_cast(source_type: &DataType, target_type: &DataType) -> bool { } /// Returns `true` when the cast from `source_type` to `target_type` is -/// value-preserving (injective) and order-preserving at the family level — -/// i.e. every value in the source domain can be round-tripped through the cast -/// without information loss, and comparisons keep the same ordering. +/// value-preserving (injective) and order-preserving at the family level. +/// Timestamp widening is classified family-safe but retains its pre-existing +/// overflow limitation at extreme source values; it is not a full-domain +/// guarantee. Other accepted casts can be round-tripped without information +/// loss and preserve comparison ordering. /// /// This is a conservative gate used by [`exact_preimage_cast`]. Timestamp /// precision narrowing is not exact-safe (it is handled by /// [`CastPredicatePreimage::Range`]); timestamp widening / equal-unit casts are -/// family-safe but still require the caller's literal round-trip check. +/// family-safe subject to the widening overflow limitation and still require the +/// caller's literal round-trip check. fn is_exact_cast_safe(source_type: &DataType, target_type: &DataType) -> bool { // Unwrap at most one level of dictionary for family-level checking. let (src, tgt) = match (source_type, target_type) { @@ -434,12 +438,17 @@ fn timestamp_narrowing_range_preimage( ) -> Result> { let ( DataType::Timestamp(source_unit, source_tz), - DataType::Timestamp(target_unit, _), + DataType::Timestamp(target_unit, target_tz), ) = (source_type, target_type) else { return Ok(None); }; + // A naive source cannot be inverted into a timezone-aware target bucket. + if source_tz.is_none() && target_tz.is_some() { + return Ok(None); + } + let source_scale = i128::from(timestamp_unit_scale(source_unit)); let target_scale = i128::from(timestamp_unit_scale(target_unit)); if source_scale <= target_scale { @@ -1772,6 +1781,14 @@ mod tests { } } + #[test] + fn test_timestamp_narrowing_range_rejects_naive_to_timezone() { + let source = DataType::Timestamp(TimeUnit::Nanosecond, None); + let target = DataType::Timestamp(TimeUnit::Millisecond, Some("+01:00".into())); + let literal = ScalarValue::TimestampMillisecond(Some(0), Some("+01:00".into())); + assert_preimage_none(&source, &target, Operator::Eq, &literal); + } + #[test] fn test_cast_predicate_preimage_timestamp_widening_ordered() { let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index 0ca5597a69fab..8d86bb9ce7d87 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -98,10 +98,24 @@ pub(super) fn supports_cast_predicate_for_binary( return false; }; - cast_predicate_preimage(&source_type, target_type, op, lit_value) - .ok() - .flatten() - .is_some() + let Ok(preimage) = cast_predicate_preimage(&source_type, target_type, op, lit_value) + else { + return false; + }; + if preimage.is_none() { + return false; + } + // Equality-like range predicates duplicate their input; volatile expressions + // cannot be duplicated. + !(matches!(preimage, Some(CastPredicatePreimage::Range(_))) + && matches!( + op, + Operator::Eq + | Operator::NotEq + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom + ) + && inner_expr.is_volatile()) } pub(super) fn supports_cast_predicate_for_inlist( diff --git a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs index 776eec2d110f6..6f9b029ecd7af 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -44,6 +44,7 @@ use crate::PhysicalExpr; use crate::expressions::{ BinaryExpr, CastExpr, Literal, TryCastExpr, is_not_null, is_null, lit, }; +use datafusion_physical_expr_common::physical_expr::is_volatile; /// Attempts to unwrap casts in comparison expressions. pub(crate) fn unwrap_cast_in_comparison( @@ -142,6 +143,18 @@ fn try_unwrap_cast_comparison( Ok(Some(Arc::new(binary_expr))) } Some(CastPredicatePreimage::Range(interval)) => { + // Equality-like range predicates duplicate their input; do not duplicate + // volatile expressions. + if matches!( + op, + Operator::Eq + | Operator::NotEq + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom + ) && is_volatile(&inner_expr) + { + return Ok(None); + } rewrite_with_preimage(interval, op, inner_expr).map(Some) } None => Ok(None), From 6da023e725f2c9bc0bc44793fa815d768469cb31 Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Thu, 8 Oct 2026 22:36:03 +0800 Subject: [PATCH 14/14] refactor: address cast-preimage review suggestions Deprecate the retained date-narrowing helper for 56.0.0 and keep its existing behavior covered. Extract input-duplication metadata onto CastPredicatePreimage and use it in the logical volatility guard. Preserve existing cast comparison, timezone and NULL semantics without changing physical rewrite logic or the supported cast allowlist. Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- datafusion/expr-common/src/casts.rs | 51 +++++++++++++++++++ .../src/simplify_expressions/cast_preimage.rs | 16 +----- 2 files changed, 53 insertions(+), 14 deletions(-) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 51db64075db6e..a0027706b5863 100644 --- a/datafusion/expr-common/src/casts.rs +++ b/datafusion/expr-common/src/casts.rs @@ -51,6 +51,21 @@ pub enum CastPredicatePreimage { Range(Interval), } +impl CastPredicatePreimage { + /// Returns whether this preimage and comparison operator produce multiple + /// predicates that each evaluate the input expression. + pub fn duplicates_input(&self, op: Operator) -> bool { + matches!(self, Self::Range(_)) + && matches!( + op, + Operator::Eq + | Operator::NotEq + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom + ) + } +} + /// Convert a literal [`ScalarValue`] to `target_type`, preserving the exact value. /// /// Returns `None` if the value cannot be represented in `target_type` @@ -678,6 +693,10 @@ fn is_lossy_temporal_cast(from_type: &DataType, to_type: &DataType) -> bool { /// This pair is already rejected by the `is_exact_cast_safe` allowlist /// (including when wrapped in a `Dictionary`), so predicate rewrites do not need /// a separate `Date64 -> Date32` early return. +#[deprecated( + since = "56.0.0", + note = "Date64 -> Date32 is rejected by the cast_predicate_preimage allowlist" +)] pub fn is_date_narrowing_cast(from_type: &DataType, to_type: &DataType) -> bool { matches!((from_type, to_type), (DataType::Date64, DataType::Date32)) } @@ -1525,6 +1544,38 @@ mod tests { } #[test] + fn test_cast_predicate_preimage_duplicates_input() { + let exact = CastPredicatePreimage::Exact(ScalarValue::Int32(Some(1))); + let range = CastPredicatePreimage::Range( + Interval::try_new(ScalarValue::Int32(Some(0)), ScalarValue::Int32(Some(2))) + .unwrap(), + ); + let cases = [ + (Operator::Eq, true), + (Operator::NotEq, true), + (Operator::IsDistinctFrom, true), + (Operator::IsNotDistinctFrom, true), + (Operator::Lt, false), + (Operator::LtEq, false), + (Operator::Gt, false), + (Operator::GtEq, false), + ]; + + for (op, range_duplicates) in cases { + assert!(!exact.duplicates_input(op), "Exact with {op:?}"); + assert_eq!( + range.duplicates_input(op), + range_duplicates, + "Range with {op:?}" + ); + } + } + + #[test] + #[expect( + deprecated, + reason = "kept to verify the deprecated helper's semantics" + )] fn test_is_date_narrowing_cast() { // Only Date64 -> Date32 narrows (ms -> days, many-to-one). assert!(is_date_narrowing_cast(&DataType::Date64, &DataType::Date32)); diff --git a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs index 8d86bb9ce7d87..f03df39562455 100644 --- a/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -102,20 +102,8 @@ pub(super) fn supports_cast_predicate_for_binary( else { return false; }; - if preimage.is_none() { - return false; - } - // Equality-like range predicates duplicate their input; volatile expressions - // cannot be duplicated. - !(matches!(preimage, Some(CastPredicatePreimage::Range(_))) - && matches!( - op, - Operator::Eq - | Operator::NotEq - | Operator::IsDistinctFrom - | Operator::IsNotDistinctFrom - ) - && inner_expr.is_volatile()) + // Volatile expressions cannot be duplicated. + preimage.is_some_and(|p| !(p.duplicates_input(op) && inner_expr.is_volatile())) } pub(super) fn supports_cast_predicate_for_inlist(