Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
354 changes: 350 additions & 4 deletions datafusion/core/tests/expr_api/simplification.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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:
///
Expand Down Expand Up @@ -152,6 +159,345 @@ fn to_timestamp_expr(arg: impl Into<String>) -> 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<dyn datafusion_physical_expr::PhysicalExpr>,
batch: &RecordBatch,
) -> Vec<Option<bool>> {
let value = physical_expr.evaluate(batch).unwrap();
let array = value.into_array(batch.num_rows()).unwrap();
array
.as_any()
.downcast_ref::<BooleanArray>()
.unwrap()
.iter()
.collect()
}

fn evaluate_boolean_expr(expr: &Expr, batch: &RecordBatch) -> Vec<Option<bool>> {
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<Option<bool>> {
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<Option<i64>> {
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<Option<i64>>,
) -> 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<Option<bool>> =
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<AtomicUsize>) -> Arc<ScalarUDF> {
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()
Expand Down
Loading
Loading