diff --git a/datafusion/core/tests/expr_api/simplification.rs b/datafusion/core/tests/expr_api/simplification.rs index 64ff1be99fbca..fd4cc872ad9c3 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, in_list, 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,345 @@ 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) +} + +/// Raw timestamp values clustered near the epoch, expressed in `unit`. They are +/// far below the one-hour `+01:00` shift, so the tz-aware target never matches +/// any non-NULL row. +fn timezone_cast_values(unit: &arrow::datatypes::TimeUnit) -> Vec> { + match unit { + arrow::datatypes::TimeUnit::Nanosecond => vec![ + Some(-1_000_000), + Some(-999_999), + Some(0), + Some(999_999), + Some(1_000_000), + None, + ], + _ => vec![ + Some(-1000), + Some(-999), + Some(0), + Some(999), + Some(1000), + None, + ], + } +} + +fn timezone_cast_array( + unit: &arrow::datatypes::TimeUnit, + values: Vec>, +) -> ArrayRef { + match unit { + arrow::datatypes::TimeUnit::Nanosecond => { + Arc::new(TimestampNanosecondArray::from(values)) + } + arrow::datatypes::TimeUnit::Millisecond => { + Arc::new(arrow::array::TimestampMillisecondArray::from(values)) + } + other => unreachable!("unsupported test time unit {other:?}"), + } +} + +fn timezone_cast_literal(unit: &arrow::datatypes::TimeUnit, value: i64) -> ScalarValue { + match unit { + arrow::datatypes::TimeUnit::Nanosecond => { + ScalarValue::TimestampNanosecond(Some(value), Some("+01:00".into())) + } + arrow::datatypes::TimeUnit::Millisecond => { + ScalarValue::TimestampMillisecond(Some(value), Some("+01:00".into())) + } + other => unreachable!("unsupported test time unit {other:?}"), + } +} + +#[test] +fn timestamp_timezone_cast_preimage_preserves_results() { + use arrow::datatypes::TimeUnit; + + // Same unit, widening and narrowing naive -> tz-aware casts: none of them may + // be preimage-rewritten, because Arrow shifts the naive value to UTC when the + // target carries a timezone. + for (source_unit, target_unit) in [ + (TimeUnit::Nanosecond, TimeUnit::Nanosecond), + (TimeUnit::Millisecond, TimeUnit::Nanosecond), + (TimeUnit::Nanosecond, TimeUnit::Millisecond), + ] { + let source_type = DataType::Timestamp(source_unit, None); + let target_type = DataType::Timestamp(target_unit, Some("+01:00".into())); + let values = timezone_cast_values(&source_unit); + let batch = RecordBatch::try_from_iter(vec![( + "ts", + timezone_cast_array(&source_unit, values.clone()), + )]) + .unwrap(); + assert_eq!(batch.schema().field(0).data_type(), &source_type); + let expected: Vec> = + values.iter().map(|value| value.map(|_| false)).collect(); + + for try_cast in [false, true] { + let make_expr = || { + cast_or_try_cast(col("ts"), target_type.clone(), try_cast) + .eq(lit(timezone_cast_literal(&target_unit, 0))) + }; + 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 preimage. + assert_eq!(evaluate_boolean_expr(&original, &batch), expected); + assert_eq!(logical_rows, expected); + assert_eq!(physical_rows, expected); + assert!( + format!("{logical}").contains("CAST(ts AS"), + "cast must be retained for {source_unit:?} -> {target_unit:?}, \ + try_cast={try_cast}: {logical}" + ); + } + + // The IN-list rewrite uses the same exact-preimage helper, so a + // multi-item list must not fold either. + let in_list_expr = in_list( + cast_or_try_cast(col("ts"), target_type.clone(), false), + vec![ + lit(timezone_cast_literal(&target_unit, 0)), + lit(timezone_cast_literal(&target_unit, 1)), + ], + false, + ); + let logical = simplify_logical_expr( + in_list_expr.clone(), + batch.schema().to_dfschema_ref().unwrap(), + ); + assert_eq!(evaluate_boolean_expr(&in_list_expr, &batch), expected); + assert_eq!(evaluate_boolean_expr(&logical, &batch), expected); + assert!( + format!("{logical}").contains("CAST(ts AS"), + "cast must be retained in the IN-list for {source_unit:?} -> \ + {target_unit:?}: {logical}" + ); + } +} + +#[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/core/tests/parquet/row_group_pruning.rs b/datafusion/core/tests/parquet/row_group_pruning.rs index c15df47f624a8..1b5527256f815 100644 --- a/datafusion/core/tests/parquet/row_group_pruning.rs +++ b/datafusion/core/tests/parquet/row_group_pruning.rs @@ -1003,6 +1003,8 @@ async fn prune_decimal_eq() { .with_expected_rows(2) .test_row_group_prune() .await; + // Increasing scale here narrows integer capacity, so the retained cast + // has no literal guarantee and its Bloom filter is not evaluated. RowGroupPruningTest::new() .with_scenario(Scenario::DecimalLargePrecision) .with_query("SELECT * FROM t where decimal_col = 4.00") @@ -1010,6 +1012,19 @@ async fn prune_decimal_eq() { .with_matched_by_stats(Some(2)) .with_pruned_by_stats(Some(1)) .with_pruned_files(Some(0)) + .with_matched_by_bloom_filter(Some(0)) + .with_pruned_by_bloom_filter(Some(0)) + .with_expected_rows(2) + .test_row_group_prune() + .await; + // A source-typed decimal literal exposes a usable Bloom-filter guarantee. + RowGroupPruningTest::new() + .with_scenario(Scenario::DecimalLargePrecision) + .with_query("SELECT * FROM t where decimal_col = cast(4.00 as decimal(38,2))") + .with_expected_errors(Some(0)) + .with_matched_by_stats(Some(2)) + .with_pruned_by_stats(Some(1)) + .with_pruned_files(Some(0)) .with_matched_by_bloom_filter(Some(2)) .with_pruned_by_bloom_filter(Some(0)) .with_expected_rows(2) diff --git a/datafusion/expr-common/src/casts.rs b/datafusion/expr-common/src/casts.rs index 3518c02772672..9fa4fead9698d 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,36 @@ 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 { + /// 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. + 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. /// @@ -76,6 +109,530 @@ 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 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( + source_type: &DataType, + target_type: &DataType, + op: Operator, + lit_value: &ScalarValue, +) -> Result> { + // Arrow adjusts a timezone-naive value to UTC when the target is + // timezone-aware (`adjust_timestamp_to_timezone`). The precision-only + // arithmetic here does not model that adjustment, so retain the cast. This + // dispatcher check covers the range, exact and widening paths below. + if is_naive_to_tz_timestamp_cast(source_type, target_type) { + return Ok(None); + } + + if let Some(preimage) = maybe_range_preimage(source_type, target_type, lit_value)? { + return Ok(Some(preimage)); + } + + if let Some(value) = + exact_preimage_int_to_str_eq_like(source_type, target_type, op, lit_value) + { + return Ok(Some(CastPredicatePreimage::Exact(value))); + } + + 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( + source_type: &DataType, + target_type: &DataType, + lit_value: &ScalarValue, +) -> Result> { + if !is_timestamp_precision_narrowing_cast(source_type, target_type) { + return Ok(None); + } + + Ok( + timestamp_narrowing_range_preimage(source_type, target_type, lit_value)? + .map(CastPredicatePreimage::Range), + ) +} + +/// 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 exact_preimage_cast( + source_type: &DataType, + target_type: &DataType, + lit_value: &ScalarValue, +) -> Option { + // 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; widening retains its pre-existing overflow + // limitation at extreme source values. + if !is_exact_cast_safe(source_type, target_type) { + return None; + } + + 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` +/// 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 is_timestamp_cast(source_type: &DataType, target_type: &DataType) -> bool { + matches!( + (source_type, target_type), + (DataType::Timestamp(_, _), DataType::Timestamp(_, _)) + ) +} + +/// Returns true when casting a timezone-naive timestamp to a timezone-aware one. +/// +/// Arrow adjusts such values to UTC (`adjust_timestamp_to_timezone`), which the +/// precision-only preimage arithmetic here does not model, so the cast is kept. +fn is_naive_to_tz_timestamp_cast(source_type: &DataType, target_type: &DataType) -> bool { + matches!( + (source_type, target_type), + ( + DataType::Timestamp(_, None), + DataType::Timestamp(_, Some(_)) + ) + ) +} + +/// Returns `true` when the cast from `source_type` to `target_type` is +/// 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 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) { + (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) { + // Timestamp precision narrowing is handled as a Range preimage, and a + // naive → timezone-aware cast reads the raw value as local wall-clock + // time and shifts it to UTC; neither is an exact preimage. + return !is_timestamp_precision_narrowing_cast(src, tgt) + && !is_naive_to_tz_timestamp_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 +} + +/// 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 exact_preimage_int_to_str_eq_like( + 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 + | Operator::IsDistinctFrom + | Operator::IsNotDistinctFrom, + 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, + } +} + +/// Computes the source-domain preimage interval for timestamp precision +/// narrowing casts. +/// +/// 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. +/// Timezone-naive sources with timezone-aware targets are rejected by +/// [`is_naive_to_tz_timestamp_cast`] in the dispatcher. +fn timestamp_narrowing_range_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) +} + +/// 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`. +/// +/// 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; + 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)) + } + } +} + +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,48 +681,25 @@ 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). /// -/// Like [`is_timestamp_precision_narrowing_cast`], this guards comparison cast -/// unwrapping against a many-to-one column cast. `CAST(date64 AS Date32) = lit_day` -/// matches any millisecond within that day, but the rewritten `date64 = lit_ms` -/// matches only midnight. Arrow does not require `Date64` values to be whole days -/// (see arrow-rs#5288), so the column may carry sub-day values the planner cannot -/// see; the widening direction (`Date32 -> Date64`) is injective and stays allowed. +/// `CAST(date64 AS Date32) = lit_day` matches any millisecond within that day, +/// but the rewritten `date64 = lit_ms` matches only midnight. Arrow does not +/// require `Date64` values to be whole days (see arrow-rs#5288), so the column +/// may carry sub-day values the planner cannot see; the widening direction +/// (`Date32 -> Date64`) is injective and stays allowed. +/// +/// 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)) } - -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 +1030,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), }; @@ -1016,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)); @@ -1253,6 +1813,240 @@ mod tests { assert_eq!(new_scalar, ScalarValue::TimestampMillisecond(None, None)); } + #[test] + fn test_cast_predicate_preimage_exact() { + assert_preimage_exact( + &DataType::Int32, + &DataType::Int64, + Operator::Gt, + &ScalarValue::Int64(Some(10)), + ScalarValue::Int32(Some(10)), + ); + + assert_preimage_exact( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("123".to_string())), + ScalarValue::Int32(Some(123)), + ); + + assert_preimage_none( + &DataType::Int32, + &DataType::Utf8, + Operator::Eq, + &ScalarValue::Utf8(Some("0123".to_string())), + ); + } + + #[test] + fn test_cast_predicate_preimage_timestamp_narrowing_range() { + 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] + 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_naive_to_timezone_timestamp_gate_rejected() { + // Same unit, widening and narrowing targets: Arrow adjusts naive + // values to UTC, which the precision-only arithmetic does not model. + // The literal is unit-aligned (0), so the rejection comes from the + // timezone adjustment and not from an unaligned widening bound. + for (source_unit, target_unit) in [ + (TimeUnit::Nanosecond, TimeUnit::Nanosecond), + (TimeUnit::Millisecond, TimeUnit::Nanosecond), + (TimeUnit::Nanosecond, TimeUnit::Millisecond), + ] { + let source = DataType::Timestamp(source_unit, None); + let target = DataType::Timestamp(target_unit, Some("+07:00".into())); + let literal = timestamp_scalar(&target_unit, Some("+07:00".into()), 0); + + assert!( + !is_exact_cast_safe(&source, &target), + "{source:?} -> {target:?} must not be exact-safe" + ); + assert!( + exact_preimage_cast(&source, &target, &literal).is_none(), + "{source:?} -> {target:?} must not have an exact preimage" + ); + for op in [ + Operator::Eq, + Operator::Gt, + Operator::LtEq, + Operator::IsDistinctFrom, + ] { + assert_preimage_none(&source, &target, op, &literal); + } + } + + // Conservative for every timezone, including fixed zero offsets: + // `Some("UTC")` is rejected like `Some("+07:00")`. + let source = DataType::Timestamp(TimeUnit::Nanosecond, None); + for tz in ["UTC", "+00:00"] { + let target = DataType::Timestamp(TimeUnit::Nanosecond, Some(tz.into())); + assert!(!is_exact_cast_safe(&source, &target)); + assert_preimage_none( + &source, + &target, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(0), Some(tz.into())), + ); + } + + // Positive control: a timezone-aware source keeps the existing accepted + // policies (same unit and widening stay exact-safe). + let aware_ms = DataType::Timestamp(TimeUnit::Millisecond, Some("+07:00".into())); + let aware_ns = DataType::Timestamp(TimeUnit::Nanosecond, Some("+07:00".into())); + assert!(is_exact_cast_safe(&aware_ns, &aware_ns)); + assert!(is_exact_cast_safe(&aware_ms, &aware_ns)); + } + + #[test] + fn test_naive_to_timezone_dictionary_wrapped_gate_rejected() { + let dict = |inner: DataType| { + DataType::Dictionary(Box::new(DataType::Int32), Box::new(inner)) + }; + let naive_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + let aware_ns = DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into())); + + // `is_exact_cast_safe` unwraps one level of dictionary, so a wrapped + // naive -> aware pair is rejected like the bare pair. + assert!(!is_exact_cast_safe(&naive_ns, &aware_ns)); + assert!(!is_exact_cast_safe( + &dict(naive_ns.clone()), + &dict(aware_ns.clone()) + )); + assert!(!is_exact_cast_safe(&dict(naive_ns.clone()), &aware_ns)); + assert!(!is_exact_cast_safe(&naive_ns, &dict(aware_ns.clone()))); + + // Safe controls stay allowed: aware -> aware unit widening, identity, + // and dropping the timezone to a naive target. + assert!(is_exact_cast_safe( + &DataType::Timestamp(TimeUnit::Millisecond, Some("UTC".into())), + &DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into())), + )); + assert!(is_exact_cast_safe(&naive_ns, &naive_ns)); + assert!(is_exact_cast_safe( + &DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into())), + &naive_ns, + )); + } + + #[test] + fn test_cast_predicate_preimage_timestamp_widening_ordered() { + let ts_ms = DataType::Timestamp(TimeUnit::Millisecond, None); + let ts_ns = DataType::Timestamp(TimeUnit::Nanosecond, None); + + assert_preimage_exact( + &ts_ms, + &ts_ns, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_000_000), None), + 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, + Operator::Eq, + &ScalarValue::TimestampNanosecond(Some(123_456_789), None), + ); + } + + #[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); + 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![ @@ -1688,4 +2482,385 @@ 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); + + // These millisecond values expand beyond i64 range in nanoseconds, + // so the preimage should return None rather than panicking. + for value in [i64::MAX, i64::MIN] { + assert_preimage_none( + &ts_ns, + &ts_ms, + Operator::Eq, + &ScalarValue::TimestampMillisecond(Some(value), 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:?}"), + } + } + + // ── 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)), + ); + // Increasing scale while narrowing integer capacity can overflow for + // valid source values, even when this particular literal round-trips. + let source_type = DataType::Decimal128(38, 2); + let target_type = DataType::Decimal128(38, 15); + let target_literal = ScalarValue::Decimal128(Some(4 * 10_i128.pow(15)), 38, 15); + assert_eq!( + try_cast_literal_to_type(&target_literal, &source_type), + Some(ScalarValue::Decimal128(Some(400), 38, 2)), + ); + assert_preimage_none(&source_type, &target_type, Operator::Eq, &target_literal); + let overflowing_source = ScalarValue::Decimal128(Some(10_i128.pow(37)), 38, 2); + let array = overflowing_source.to_array_of_size(1).unwrap(); + assert!( + cast_with_options( + &array, + &target_type, + &CastOptions { + safe: false, + ..Default::default() + }, + ) + .is_err() + ); + + // 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:?}" + ); + } + + 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 + i128::from(target_value.rem_euclid(quotient) != 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); + 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 new file mode 100644 index 0000000000000..f03df39562455 --- /dev/null +++ b/datafusion/optimizer/src/simplify_expressions/cast_preimage.rs @@ -0,0 +1,565 @@ +// 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_preimage, exact_preimage_cast, +}; + +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; + }; + + let Ok(preimage) = cast_predicate_preimage(&source_type, target_type, op, lit_value) + else { + return false; + }; + // 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( + 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; + }; + + // 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, _) => { + exact_preimage_cast(&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) = + exact_preimage_cast(&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::{binary_expr, cast, col, in_list, try_cast}; + + #[test] + fn test_exact_preimage_cast_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(); + + for (op, lit_ms, expected) in [ + ( + 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(); + + for (op, expected) in [ + ( + 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] + fn test_cast_preimage_timestamp_widening_ordered() { + 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); + + 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 = 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(); + + for (op, expected) in [ + ( + 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] + 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() + .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("ts_milli", timestamp_millis_type(), 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 6d75eb4352bb4..009d65c0f26ed 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}; @@ -1990,28 +1988,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, @@ -2021,52 +2018,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/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 3e67fc8291a4e..f03c65e406820 100644 --- a/datafusion/physical-expr/src/simplifier/unwrap_cast.rs +++ b/datafusion/physical-expr/src/simplifier/unwrap_cast.rs @@ -34,15 +34,15 @@ use std::sync::Arc; use arrow::datatypes::{DataType, Schema}; -use datafusion_common::{Result, ScalarValue, tree_node::Transformed}; +use datafusion_common::{Result, ScalarValue, internal_err, 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, -}; +use datafusion_expr_common::casts::{CastPredicatePreimage, cast_predicate_preimage}; use crate::PhysicalExpr; -use crate::expressions::{BinaryExpr, CastExpr, Literal, TryCastExpr, lit}; +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( @@ -69,8 +69,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 +88,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 +122,92 @@ 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) - { - return Ok(None); + 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)) => { + // 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), } +} - // 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))); - } +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)?, + ), + _ => return internal_err!("Expect comparison operators, got {op}"), + }; + + Ok(rewritten_expr) +} - Ok(None) +fn binary( + left: Arc, + op: Operator, + right: Arc, +) -> Arc { + Arc::new(BinaryExpr::new(left, op, right)) } #[cfg(test)] @@ -180,6 +244,47 @@ 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))) + } + + 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(); @@ -251,6 +356,100 @@ mod tests { assert!(!result.transformed); } + /// Unit-specific timestamp literal, e.g. `0` in milliseconds or nanoseconds. + fn timestamp_literal( + unit: &TimeUnit, + value: i64, + tz: Option>, + ) -> ScalarValue { + match unit { + TimeUnit::Nanosecond => ScalarValue::TimestampNanosecond(Some(value), tz), + TimeUnit::Millisecond => ScalarValue::TimestampMillisecond(Some(value), tz), + other => unreachable!("unsupported test time unit {other:?}"), + } + } + + #[test] + fn test_no_unwrap_naive_to_timezone_timestamp() { + let tz: Option> = Some("+07:00".into()); + + // cast(naive_ts AS Timestamp(unit, "+07:00")) must NOT unwrap for CAST or + // TRY_CAST: Arrow adjusts naive values to UTC, which the precision-only + // preimage arithmetic does not model. Same unit, widening and narrowing + // targets are all rejected by the shared gate. + for (source_unit, target_unit) in [ + (TimeUnit::Nanosecond, TimeUnit::Nanosecond), + (TimeUnit::Millisecond, TimeUnit::Nanosecond), + (TimeUnit::Nanosecond, TimeUnit::Millisecond), + ] { + let schema = timestamp_schema(source_unit); + let column_expr = col("ts", &schema).unwrap(); + let target_type = DataType::Timestamp(target_unit, tz.clone()); + let literal = lit(timestamp_literal(&target_unit, 0, tz.clone())); + + for try_cast in [false, true] { + let cast_expr: Arc = if try_cast { + Arc::new(TryCastExpr::new( + Arc::clone(&column_expr), + target_type.clone(), + )) + } else { + Arc::new(CastExpr::new( + Arc::clone(&column_expr), + target_type.clone(), + None, + )) + }; + let binary_expr = Arc::new(BinaryExpr::new( + cast_expr, + Operator::Eq, + Arc::clone(&literal), + )); + + let result = unwrap_cast_in_comparison(binary_expr, &schema).unwrap(); + assert!( + !result.transformed, + "unexpected rewrite: {source_unit:?} -> {target_unit:?}, try_cast={try_cast}" + ); + } + } + + // Safe control: the same widening to a *naive* target still unwraps. + let ms_schema = timestamp_schema(TimeUnit::Millisecond); + let column_expr = col("ts", &ms_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, &ms_schema).unwrap(); + assert!(result.transformed); + } + + #[test] + fn test_rewrite_with_preimage_rejects_non_comparison_operator() { + // Range preimages only map onto comparison operators; anything else is + // an internal error rather than a panic. + let schema = test_schema(); + let expr = col("c1", &schema).unwrap(); + let interval = datafusion_expr_common::interval_arithmetic::Interval::try_new( + ScalarValue::Int32(Some(0)), + ScalarValue::Int32(Some(10)), + ) + .unwrap(); + + let err = rewrite_with_preimage(interval, Operator::Plus, expr).unwrap_err(); + assert!( + err.to_string() + .contains("Expect comparison operators, got +"), + "unexpected error: {err}" + ); + } + #[test] fn test_no_unwrap_when_types_unsupported() { let schema = Schema::new(vec![Field::new("f1", DataType::Float32, false)]); @@ -578,55 +777,265 @@ mod tests { } #[test] - fn test_not_unwrap_timestamp_precision_narrowing() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )]); + fn test_timestamp_precision_narrowing_range_preimage_gt() { + let schema = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, + Operator::Gt, + ScalarValue::TimestampMillisecond(Some(1000), None), + ); - 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 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::TimestampNanosecond(Some(1_001_000_000), None) + ); + } + + #[test] + fn test_timestamp_precision_narrowing_range_preimage_eq() { + 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(); + assert!(result.transformed); + 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!( + 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) + ); + } + + #[test] + fn test_timestamp_widening_ordered() { + let schema = timestamp_schema(TimeUnit::Millisecond); + assert_timestamp_widening_rewrite( + &schema, + timestamp_cast_comparison( + &schema, + TimeUnit::Nanosecond, + Operator::GtEq, + ScalarValue::TimestampNanosecond(Some(123_000_000), None), + ), + Operator::GtEq, + 123, + ); + + 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 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( + try_cast_expr, + Operator::GtEq, + lit(ScalarValue::TimestampNanosecond(Some(123_456_789), None)), + )), + Operator::GtEq, + 124, + ); + + 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_unwrap_timestamp_precision_widening() { - let schema = Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Millisecond, None), - false, - )]); + 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 column_expr = col("ts", &schema).unwrap(); - let cast_expr = Arc::new(CastExpr::new( - column_expr, + let try_cast_expr = Arc::new(TryCastExpr::new( + col("ts", &schema).unwrap(), DataType::Timestamp(TimeUnit::Nanosecond, None), - None, )); - let literal_expr = lit(ScalarValue::TimestampNanosecond(Some(1_000_000), None)); - let binary_expr = - Arc::new(BinaryExpr::new(cast_expr, Operator::Eq, literal_expr)); + 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); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, + Operator::IsDistinctFrom, + ScalarValue::TimestampMillisecond(Some(1000), None), + ); 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(); + + // 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!( - right_literal.value(), - &ScalarValue::TimestampMillisecond(Some(1), None) + 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 = timestamp_schema(TimeUnit::Nanosecond); + let binary_expr = timestamp_cast_comparison( + &schema, + TimeUnit::Millisecond, + Operator::IsNotDistinctFrom, + ScalarValue::TimestampMillisecond(Some(1000), None), + ); + + 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) ); } @@ -636,6 +1045,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); @@ -654,23 +1065,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())); } } diff --git a/datafusion/sqllogictest/test_files/datetime/timestamps.slt b/datafusion/sqllogictest/test_files/datetime/timestamps.slt index 71118ddcabb73..c868bbfdbe347 100644 --- a/datafusion/sqllogictest/test_files/datetime/timestamps.slt +++ b/datafusion/sqllogictest/test_files/datetime/timestamps.slt @@ -4822,9 +4822,15 @@ SELECT column1 FROM t_utc WHERE column1 < '2024-02-01T00:00:00' AT TIME ZONE 'Am 2024-01-01T00:00:01Z 2024-02-01T00:00:01Z +# Brussels midnight is 23:00Z; Los Angeles 16:00 is 00:00Z, one hour later. query P SELECT column1 FROM t_europe WHERE column1 = '2024-01-31T16:00:01' AT TIME ZONE 'America/Los_Angeles'; ---- + +# Los Angeles 15:00 is the same instant as Brussels midnight. +query P +SELECT column1 FROM t_europe WHERE column1 = '2024-01-31T15:00:01' AT TIME ZONE 'America/Los_Angeles'; +---- 2024-02-01T00:00:01+01:00 query P diff --git a/datafusion/sqllogictest/test_files/datetime/timestamps_timezone.slt b/datafusion/sqllogictest/test_files/datetime/timestamps_timezone.slt index 42ae7c6b53602..3327b6df8d4e8 100644 --- a/datafusion/sqllogictest/test_files/datetime/timestamps_timezone.slt +++ b/datafusion/sqllogictest/test_files/datetime/timestamps_timezone.slt @@ -610,23 +610,21 @@ SELECT column1 FROM naive_col WHERE column1 = '2024-07-01T12:00:00Z'::timestampt ---- 2024-07-01T12:00:00 -# KNOWN WRONG. The same comparison under a non-UTC session zone. The literal is -# now tz-aware (18:00Z, i.e. 12:00 Denver), so coercion casts the tz-naive -# column to America/Denver, which reads its 12:00 as 12:00 Denver and should -# match the July row. unwrap_cast then moves the cast off the column and onto -# the literal, and that cast renders the literal as its UTC wall clock (18:00), -# so no row matches. PostgreSQL under TimeZone='America/Denver' returns 1. -# See https://github.com/apache/datafusion/issues/25095 +# tz-naive column compared to a tz-aware literal under a non-UTC session zone. +# The tz-aware literal is 18:00Z, i.e. 12:00 Denver, so the column's 12:00 read +# in America/Denver matches it. The cast stays on the column (it cannot move to +# the literal: Arrow adjusts naive values to UTC when the target is tz-aware). +# PostgreSQL under TimeZone='America/Denver' returns 1. statement ok SET datafusion.execution.time_zone = 'America/Denver' query I SELECT count(*) FROM naive_col WHERE column1 = '2024-07-01T18:00:00Z'::timestamptz ---- -0 +1 -# The plan shows the rewrite: the filter compares the naive column against a -# naive 18:00 (1719856800 s), not against 12:00. +# The plan keeps the cast on the naive column, comparing against the tz-aware +# 18:00 (1719856800 s) literal. statement ok set datafusion.explain.logical_plan_only = true @@ -634,7 +632,7 @@ query TT EXPLAIN SELECT column1 FROM naive_col WHERE column1 = '2024-07-01T18:00:00Z'::timestamptz ---- logical_plan -01)Filter: naive_col.column1 = TimestampNanosecond(1719856800000000000, None) +01)Filter: CAST(naive_col.column1 AS Timestamp(ns, "America/Denver")) = TimestampNanosecond(1719856800000000000, Some("America/Denver")) 02)--TableScan: naive_col projection=[column1] statement ok diff --git a/datafusion/sqllogictest/test_files/simplify_expr.slt b/datafusion/sqllogictest/test_files/simplify_expr.slt index c70ff2b955e19..eac4f91f3867e 100644 --- a/datafusion/sqllogictest/test_files/simplify_expr.slt +++ b/datafusion/sqllogictest/test_files/simplify_expr.slt @@ -518,3 +518,144 @@ 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_preimage_ts.ts_milli >= TimestampMillisecond(124, None) +02)--TableScan: cast_preimage_ts projection=[ts_nano, ts_milli] +physical_plan +01)FilterExec: ts_milli@1 >= 124 +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; + +# Naive -> timezone-aware timestamp casts are not preimage-rewritable. +# +# Arrow reads a timezone-naive timestamp as local wall-clock time and shifts it +# to UTC when the target is timezone-aware (`adjust_timestamp_to_timezone`), so +# `CAST(ts AS Timestamp(unit, Some(tz))) OP literal` is not a pure unit change. +# The shared precision arithmetic works on raw source units and does not model +# that shift, so the CAST must stay on the column for same-unit, widening and +# narrowing targets alike. +# +# `2024-01-01T00:00:00.5` is 1704067200500000000 ns of naive wall-clock time; +# read in `+07:00` and shifted to UTC it is 1704042000500000000 ns +# (1704042000500 ms). The rewritten (buggy) predicate would compare the raw +# source units against those shifted values and drop the row. +statement ok +create table cast_preimage_tz(ts timestamp, l timestamp) as values + ('2024-01-01T00:00:00.5'::timestamp, '2024-01-01T00:00:00.5'::timestamp); + +statement ok +create table cast_preimage_tz_ms as +select arrow_cast('2024-01-01T00:00:00'::timestamp, 'Timestamp(Millisecond, None)') as ts_milli; + +# Column-vs-column control: there is no literal to fold, so the CAST is kept and +# the comparison is evaluated as written. +query I +select count(*) from cast_preimage_tz +where arrow_cast(ts, 'Timestamp(Millisecond, Some("+07:00"))') + = arrow_cast(l, 'Timestamp(Millisecond, Some("+07:00"))'); +---- +1 + +# Narrowing (ns -> ms) with a literal. This pair was already guarded before the +# shared gate existed, so it stays at 1; without the guard the rewrite would +# return 0. +query I +select count(*) from cast_preimage_tz +where arrow_cast(ts, 'Timestamp(Millisecond, Some("+07:00"))') + = arrow_cast('2024-01-01T00:00:00.5'::timestamp, 'Timestamp(Millisecond, Some("+07:00"))'); +---- +1 + +# Same unit, only the timezone changes: this is the pair the shared gate has to +# reject. Without it the literal folds to a naive nanosecond literal and the +# comparison drops the row (0). +query I +select count(*) from cast_preimage_tz +where arrow_cast(ts, 'Timestamp(Nanosecond, Some("+07:00"))') + = arrow_cast('2024-01-01T00:00:00.5'::timestamp, 'Timestamp(Nanosecond, Some("+07:00"))'); +---- +1 + +# Widening (ms -> ns) with an aligned literal goes through the exact-preimage +# path and must be rejected there too. +query I +select count(*) from cast_preimage_tz_ms +where arrow_cast(ts_milli, 'Timestamp(Nanosecond, Some("+07:00"))') + = arrow_cast('2024-01-01T00:00:00'::timestamp, 'Timestamp(Nanosecond, Some("+07:00"))'); +---- +1 + +# The IN-list rewrite uses the same exact-preimage helper and must be rejected +# as well. Two literal items keep this on the IN-list path rather than folding +# to a single equality. +query I +select count(*) from cast_preimage_tz +where arrow_cast(ts, 'Timestamp(Nanosecond, Some("+07:00"))') + in (arrow_cast('2024-01-01T00:00:00.5'::timestamp, 'Timestamp(Nanosecond, Some("+07:00"))'), + arrow_cast('2023-06-01T00:00:00'::timestamp, 'Timestamp(Nanosecond, Some("+07:00"))')); +---- +1 + +statement ok +drop table cast_preimage_tz; + +statement ok +drop table cast_preimage_tz_ms;