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
1 change: 1 addition & 0 deletions native-engine/datafusion-ext-functions/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@ pub fn create_auron_ext_function(
"Spark_WeekOfYear" => shared_function!(spark_dates::spark_weekofyear),
"Spark_Quarter" => shared_function!(spark_dates::spark_quarter),
"Spark_LastDay" => shared_function!(spark_dates::spark_last_day),
"Spark_DateDiff" => shared_function!(spark_dates::spark_datediff),
"Spark_MakeDate" => shared_function!(spark_dates::spark_make_date),
"Spark_Hour" => shared_function!(spark_dates::spark_hour),
"Spark_Minute" => shared_function!(spark_dates::spark_minute),
Expand Down
65 changes: 65 additions & 0 deletions native-engine/datafusion-ext-functions/src/spark_dates.rs
Original file line number Diff line number Diff line change
Expand Up @@ -297,6 +297,29 @@ pub fn spark_last_day(args: &[ColumnarValue]) -> Result<ColumnarValue> {
Ok(ColumnarValue::Array(Arc::new(last_day)))
}

pub fn spark_datediff(args: &[ColumnarValue]) -> Result<ColumnarValue> {
let dates = ColumnarValue::values_to_arrays(args)?;
let end_date = cast(&dates[0], &DataType::Date32)?;
let start_date = cast(&dates[1], &DataType::Date32)?;
let end_date = end_date
.as_any()
.downcast_ref::<Date32Array>()
.expect("cast to Date32 must succeed");
let start_date = start_date
.as_any()
.downcast_ref::<Date32Array>()
.expect("cast to Date32 must succeed");
let result = Int32Array::from_iter(end_date.iter().zip(start_date.iter()).map(
|(end_date, start_date)| {
end_date
.zip(start_date)
.map(|(end_date, start_date)| end_date.wrapping_sub(start_date))
},
));

Ok(ColumnarValue::Array(Arc::new(result)))
}
Comment on lines +300 to +321

pub fn spark_make_date(args: &[ColumnarValue]) -> Result<ColumnarValue> {
if args.len() != 4 {
return Err(DataFusionError::Execution(
Expand Down Expand Up @@ -665,6 +688,48 @@ mod tests {
Ok(())
}

#[test]
fn test_spark_datediff() -> Result<()> {
let date = |year, month, day| {
NaiveDate::from_ymd_opt(year, month, day)
.expect("test date must be valid")
.to_epoch_days()
};
let end_date = Arc::new(Date32Array::from(vec![
Some(date(2009, 7, 31)),
Some(date(2009, 7, 30)),
Some(date(2024, 1, 1)),
Some(date(2024, 3, 1)),
Some(date(2025, 1, 1)),
None,
Some(date(2024, 1, 1)),
]));
let start_date = Arc::new(Date32Array::from(vec![
Some(date(2009, 7, 30)),
Some(date(2009, 7, 31)),
Some(date(2024, 1, 1)),
Some(date(2024, 2, 28)),
Some(date(2024, 12, 31)),
Some(date(2024, 1, 1)),
None,
]));
let args = vec![
ColumnarValue::Array(end_date),
ColumnarValue::Array(start_date),
];
let expected_ret: ArrayRef = Arc::new(Int32Array::from(vec![
Some(1),
Some(-1),
Some(0),
Some(2),
Some(1),
None,
None,
]));
assert_eq!(&spark_datediff(&args)?.into_array(7)?, &expected_ret);
Ok(())
}

#[test]
fn test_spark_make_date_null_and_invalid_inputs() -> Result<()> {
let result = spark_make_date(&[
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,24 @@ class AuronFunctionSuite extends AuronQueryTest with BaseAuronSQLSuite {
}
}

test("datediff function") {
withTable("t1") {
sql("create table t1(end_date date, start_date date) using parquet")
sql("""insert into t1 values
| (date'2009-07-31', date'2009-07-30'),
| (date'2009-07-30', date'2009-07-31'),
| (date'2024-03-01', date'2024-02-28'),
| (date'2024-01-01', date'2024-01-01'),
| (date'2025-01-01', date'2024-12-31'),
| (null, date'2024-01-01'),
| (date'2024-01-01', null)
|""".stripMargin)

checkSparkAnswerAndOperator(
"select datediff(end_date, start_date), datediff(end_date, date'2024-01-01') from t1")
}
Comment on lines +218 to +220
}

test("date-part functions with non-UTC timezone") {
withTable("t1") {
sql("create table t1(c1 timestamp) using parquet")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1002,6 +1002,8 @@ object NativeConverters extends Logging {
buildTimePartExt("Spark_Quarter", child, isPruningExpr, fallback)
case e: LastDay =>
buildExtScalarFunction("Spark_LastDay", e.children, e.dataType)
case e: DateDiff =>
buildExtScalarFunction("Spark_DateDiff", e.children, e.dataType)
Comment on lines +1005 to +1006

case e: Levenshtein =>
buildScalarFunction(pb.ScalarFunction.Levenshtein, e.children, e.dataType)
Expand Down
Loading