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
2 changes: 1 addition & 1 deletion docs/source/user-guide/latest/expressions.md
Original file line number Diff line number Diff line change
Expand Up @@ -277,7 +277,7 @@ The type-name conversion functions (`bigint`, `binary`, `boolean`, `date`, `deci
| `localtimestamp` | ✅ | — | |
| `make_date` | ✅ | Native | |
| `make_dt_interval` | ✅ | Codegen dispatch | |
| `make_interval` | ✅ | Hybrid | Routes through the JVM codegen dispatcher by default; intervals outside Arrow's nanosecond range are tracked by [#5279](https://github.com/apache/datafusion-comet/issues/5279); the native path is opt-in via allowIncompatible ([details](compatibility/expressions/datetime.md)) |
| `make_interval` | ✅ | Native | |
| `make_time` | 🔜 | — | Spark 4.1 TIME type; tracked by [#4288](https://github.com/apache/datafusion-comet/issues/4288) |
| `make_timestamp` | ✅ | Hybrid | |
| `make_timestamp_ltz` | ✅ | — | 2-arg TIME form falls back |
Expand Down
1 change: 0 additions & 1 deletion native/Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

5 changes: 2 additions & 3 deletions native/core/src/execution/serde.rs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ use datafusion_comet_proto::{
spark_expression::DataType,
spark_operator,
};
use datafusion_comet_spark_expr::calendar_interval_type;
use prost::Message;
use std::{io::Cursor, sync::Arc};

Expand Down Expand Up @@ -102,9 +103,7 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType {
// Spark's DayTimeIntervalType stores microseconds in an int64, which matches Arrow
// Duration(Microsecond) rather than the lossy Interval(DayTime) {days, millis} layout.
DataTypeId::DayTimeInterval => ArrowDataType::Duration(TimeUnit::Microsecond),
// Spark's CalendarIntervalType stores months, days, and microseconds. Arrow stores the
// same components with nanosecond precision.
DataTypeId::CalendarInterval => ArrowDataType::Interval(IntervalUnit::MonthDayNano),
DataTypeId::CalendarInterval => calendar_interval_type(),
DataTypeId::Null => ArrowDataType::Null,
DataTypeId::List => match dt_value
.type_info
Expand Down
3 changes: 1 addition & 2 deletions native/spark-expr/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,6 @@ edition = { workspace = true }
arrow = { workspace = true }
chrono = { workspace = true }
datafusion = { workspace = true }
datafusion-spark = { workspace = true }
chrono-tz = { workspace = true }
num = { workspace = true }
regex = { workspace = true }
Expand Down Expand Up @@ -222,4 +221,4 @@ harness = false

[[bench]]
name = "cast_int_to_decimal"
harness = false
harness = false
181 changes: 142 additions & 39 deletions native/spark-expr/src/datetime_funcs/make_interval.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,72 +16,175 @@
// under the License.

use crate::arithmetic_overflow_error;
use arrow::array::Array;
use arrow::datatypes::DataType;
use arrow::array::{Array, ArrayRef, Decimal128Array, Int32Array, Int64Array, StructArray};
use arrow::buffer::NullBuffer;
use arrow::datatypes::{DataType, Field, Fields};
use datafusion::common::Result;
use datafusion::logical_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature};
use datafusion_spark::function::datetime::make_interval::SparkMakeInterval as DataFusionMakeInterval;
use datafusion::logical_expr::{
ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility,
};
use std::collections::HashMap;
use std::sync::Arc;

const CALENDAR_INTERVAL_STRUCT_KEY: &str = "SPARK::calendarInterval::struct";
const MICROS_PER_HOUR: i64 = 3_600_000_000;
const MICROS_PER_MINUTE: i64 = 60_000_000;

pub fn calendar_interval_type() -> DataType {
let months = Field::new("months", DataType::Int32, false).with_metadata(HashMap::from([(
CALENDAR_INTERVAL_STRUCT_KEY.to_string(),
"true".to_string(),
)]));
DataType::Struct(Fields::from(vec![
months,
Field::new("days", DataType::Int32, false),
Field::new("microseconds", DataType::Int64, false),
]))
}

#[derive(Debug, PartialEq, Eq, Hash)]
pub struct SparkMakeInterval {
inner: DataFusionMakeInterval,
signature: Signature,
fail_on_error: bool,
}

impl SparkMakeInterval {
pub fn new(fail_on_error: bool) -> Self {
Self {
inner: DataFusionMakeInterval::new(),
signature: Signature::exact(
vec![
DataType::Int32,
DataType::Int32,
DataType::Int32,
DataType::Int32,
DataType::Int32,
DataType::Int32,
DataType::Decimal128(18, 6),
],
Volatility::Immutable,
),
fail_on_error,
}
}
}

fn make_interval(
years: i32,
months: i32,
weeks: i32,
days: i32,
hours: i32,
minutes: i32,
seconds_micros: i128,
) -> Option<(i32, i32, i64)> {
let months = years.checked_mul(12)?.checked_add(months)?;
let days = weeks.checked_mul(7)?.checked_add(days)?;
let micros = i64::try_from(seconds_micros)
.ok()?
.checked_add(i64::from(hours).checked_mul(MICROS_PER_HOUR)?)?
.checked_add(i64::from(minutes).checked_mul(MICROS_PER_MINUTE)?)?;
Some((months, days, micros))
}

impl ScalarUDFImpl for SparkMakeInterval {
fn name(&self) -> &str {
self.inner.name()
"make_interval"
}

fn signature(&self) -> &Signature {
self.inner.signature()
&self.signature
}

fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
self.inner.return_type(arg_types)
fn return_type(&self, _: &[DataType]) -> Result<DataType> {
Ok(calendar_interval_type())
}

fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let inputs = if self.fail_on_error {
Some(args.args.clone())
} else {
None
};
let result = self.inner.invoke_with_args(args)?;

if let Some(inputs) = inputs {
let inputs_are_valid = |i| {
inputs.iter().all(|input| match input {
ColumnarValue::Array(values) => values.is_valid(i),
ColumnarValue::Scalar(value) => !value.is_null(),
})
};
let overflow = match &result {
ColumnarValue::Array(values) => values.nulls().is_some_and(|nulls| {
nulls.null_count() != 0
&& nulls
.iter()
.enumerate()
.any(|(i, is_valid)| !is_valid && inputs_are_valid(i))
}),
ColumnarValue::Scalar(value) => value.is_null() && inputs_are_valid(0),
};
if overflow {
// Spark identifies the integer or long operation that overflowed. The native
// wrapper only sees the result null mask, so it can only report interval overflow.
return Err(arithmetic_overflow_error("interval").into());
let number_rows = args.number_rows;
let arrays = args
.args
.into_iter()
.map(|arg| arg.into_array(number_rows))
.collect::<Result<Vec<_>>>()?;
let years = arrays[0].as_any().downcast_ref::<Int32Array>().unwrap();
let months = arrays[1].as_any().downcast_ref::<Int32Array>().unwrap();
let weeks = arrays[2].as_any().downcast_ref::<Int32Array>().unwrap();
let days = arrays[3].as_any().downcast_ref::<Int32Array>().unwrap();
let hours = arrays[4].as_any().downcast_ref::<Int32Array>().unwrap();
let minutes = arrays[5].as_any().downcast_ref::<Int32Array>().unwrap();
let seconds = arrays[6]
.as_any()
.downcast_ref::<Decimal128Array>()
.unwrap();

let mut result_months = Vec::with_capacity(years.len());
let mut result_days = Vec::with_capacity(years.len());
let mut result_micros = Vec::with_capacity(years.len());
let mut valid = Vec::with_capacity(years.len());

for i in 0..years.len() {
if arrays.iter().any(|array| array.is_null(i)) {
result_months.push(0);
result_days.push(0);
result_micros.push(0);
valid.push(false);
continue;
}

match make_interval(
years.value(i),
months.value(i),
weeks.value(i),
days.value(i),
hours.value(i),
minutes.value(i),
seconds.value(i),
) {
Some((months, days, micros)) => {
result_months.push(months);
result_days.push(days);
result_micros.push(micros);
valid.push(true);
}
None if self.fail_on_error => {
return Err(arithmetic_overflow_error("interval").into());
}
None => {
result_months.push(0);
result_days.push(0);
result_micros.push(0);
valid.push(false);
}
}
}

Ok(result)
let columns: Vec<ArrayRef> = vec![
Arc::new(Int32Array::from(result_months)),
Arc::new(Int32Array::from(result_days)),
Arc::new(Int64Array::from(result_micros)),
];
let DataType::Struct(fields) = calendar_interval_type() else {
unreachable!()
};
Ok(ColumnarValue::Array(Arc::new(StructArray::new(
fields,
columns,
Some(NullBuffer::from(valid)),
))))
}
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn preserves_spark_microsecond_range_and_overflow() {
assert_eq!(
make_interval(1, 2, 3, 4, 2_562_048, 0, 123_456_789_012_123_456),
Some((14, 25, 132_680_161_812_123_456))
);
assert!(make_interval(i32::MAX, 0, 0, 0, 0, 0, 0).is_none());
assert!(make_interval(0, 0, 0, 0, i32::MAX, i32::MAX, i128::MAX).is_none());
}
}
2 changes: 1 addition & 1 deletion native/spark-expr/src/datetime_funcs/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ pub use extract_date_part::SparkMinute;
pub use extract_date_part::SparkSecond;
pub use hours::SparkHoursTransform;
pub use make_date::SparkMakeDate;
pub use make_interval::SparkMakeInterval;
pub use make_interval::{calendar_interval_type, SparkMakeInterval};
pub use make_time::SparkMakeTime;
pub use next_day::SparkNextDay;
pub use seconds_to_timestamp::SparkSecondsToTimestamp;
Expand Down
8 changes: 4 additions & 4 deletions native/spark-expr/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -76,10 +76,10 @@ pub use comet_scalar_funcs::{
};
pub use csv_funcs::*;
pub use datetime_funcs::{
spark_day_name, spark_month_name, spark_to_time, SparkDateDiff, SparkDateFromUnixDate,
SparkDateTrunc, SparkHour, SparkHoursTransform, SparkMakeDate, SparkMakeInterval,
SparkMakeTime, SparkMinute, SparkNextDay, SparkSecond, SparkSecondsToTimestamp,
SparkUnixTimestamp, TimestampTruncExpr,
calendar_interval_type, spark_day_name, spark_month_name, spark_to_time, SparkDateDiff,
SparkDateFromUnixDate, SparkDateTrunc, SparkHour, SparkHoursTransform, SparkMakeDate,
SparkMakeInterval, SparkMakeTime, SparkMinute, SparkNextDay, SparkSecond,
SparkSecondsToTimestamp, SparkUnixTimestamp, TimestampTruncExpr,
};
pub use error::{decimal_overflow_error, SparkError, SparkErrorWithContext, SparkResult};
pub use hash_funcs::*;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ import org.apache.spark.sql.types._

import org.apache.comet.codegen.CometBatchKernelCodegen.{ArrayColumnSpec, ArrowColumnSpec, MapColumnSpec, ScalarColumnSpec, StructColumnSpec}
import org.apache.comet.shims.CometTypeShim
import org.apache.comet.vector.CometPlainVector
import org.apache.comet.vector.{CometPlainVector, CometStructVector}

/**
* Input-side emitters for the codegen kernel: typed field declarations, per-batch input casts,
Expand Down Expand Up @@ -69,6 +69,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
classOf[IntervalYearVector],
classOf[IntervalMonthDayNanoVector])
private val cometPlainVectorName: String = classOf[CometPlainVector].getName
private val cometStructVectorName: String = classOf[CometStructVector].getName

/** Emit kernel typed-vector field declarations for every level of every input column. */
def emitInputFieldDecls(inputSchema: Seq[ArrowColumnSpec]): String = {
Expand Down Expand Up @@ -109,11 +110,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
if (!spec.nullable) {
s" case $ord: return false;"
} else {
// CometPlainVector exposes `isNullAt`; Arrow-typed fields expose `isNull`. Same semantics.
val method = spec.vectorClass match {
case cls if wrapsInCometPlainVector(cls) => "isNullAt"
case _ => "isNull"
}
val method = nullCheckMethod(spec)
s" case $ord: return this.col$ord.$method(this.rowIdx);"
}
}
Expand Down Expand Up @@ -146,7 +143,8 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
s" case $ord: return this.col$ord.getLong(this.rowIdx);"
}
val intervalCases = withOrd.collect {
case (ArrowColumnSpec(cls, _), ord) if cls == classOf[IntervalMonthDayNanoVector] =>
case (ArrowColumnSpec(cls, _), ord)
if cls == classOf[IntervalMonthDayNanoVector] || cls == classOf[StructVector] =>
s" case $ord: return this.col$ord.getInterval(this.rowIdx);"
}
val floatCases = withOrd.collect {
Expand Down Expand Up @@ -425,6 +423,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
*/
def nullCheckMethod(spec: ArrowColumnSpec): String = spec match {
case sc: ScalarColumnSpec if wrapsInCometPlainVector(sc.vectorClass) => "isNullAt"
case sc: ScalarColumnSpec if sc.vectorClass == classOf[StructVector] => "isNullAt"
case _ => "isNull"
}

Expand All @@ -436,9 +435,13 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
// Primitive scalars wrap in CometPlainVector for JIT-inlined Platform.get* against a
// cached buffer address. Decimal/VarChar/VarBinary stay on the Arrow typed field with
// cached data- (and offset-) buffer addresses for inline unsafe reads.
val fieldClass =
if (wrapsInCometPlainVector(sc.vectorClass)) cometPlainVectorName
else sc.vectorClass.getName
val fieldClass = if (wrapsInCometPlainVector(sc.vectorClass)) {
cometPlainVectorName
} else if (sc.vectorClass == classOf[StructVector]) {
cometStructVectorName
} else {
sc.vectorClass.getName
}
out += s"private $fieldClass $path;"
if (needsValueAddrField(sc.vectorClass)) {
out += s"private long ${path}_valueAddr;"
Expand Down Expand Up @@ -471,6 +474,8 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim {
case sc: ScalarColumnSpec =>
if (wrapsInCometPlainVector(sc.vectorClass)) {
out += s"this.$path = new $cometPlainVectorName($source);"
} else if (sc.vectorClass == classOf[StructVector]) {
out += s"this.$path = new $cometStructVectorName($source, null);"
} else {
out += s"this.$path = (${sc.vectorClass.getName}) $source;"
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim {
case TimestampNTZType => classOf[TimeStampMicroVector].getName
case _: YearMonthIntervalType => classOf[IntervalYearVector].getName
case _: DayTimeIntervalType => classOf[DurationVector].getName
case CalendarIntervalType => classOf[IntervalMonthDayNanoVector].getName
case CalendarIntervalType => classOf[StructVector].getName
case _: ArrayType => classOf[ListVector].getName
case _: StructType => classOf[StructVector].getName
case _: MapType => classOf[MapVector].getName
Expand Down Expand Up @@ -222,11 +222,22 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim {
case CalendarIntervalType =>
val set = if (nested) "setSafe" else "set"
val interval = ctx.freshName("interval")
val months = ctx.freshName("intervalMonths")
val days = ctx.freshName("intervalDays")
val micros = ctx.freshName("intervalMicros")
OutputEmit(
"",
s"""${classOf[IntVector].getName} $months =
| (${classOf[IntVector].getName}) $targetVec.getChildByOrdinal(0);
|${classOf[IntVector].getName} $days =
| (${classOf[IntVector].getName}) $targetVec.getChildByOrdinal(1);
|${classOf[BigIntVector].getName} $micros =
| (${classOf[
BigIntVector].getName}) $targetVec.getChildByOrdinal(2);""".stripMargin,
s"""org.apache.spark.unsafe.types.CalendarInterval $interval = $source;
|$targetVec.$set($idx, $interval.months, $interval.days,
| java.lang.Math.multiplyExact($interval.microseconds, 1000L));""".stripMargin)
|$targetVec.setIndexDefined($idx);
|$months.$set($idx, $interval.months);
|$days.$set($idx, $interval.days);
|$micros.$set($idx, $interval.microseconds);""".stripMargin)
case dt if isTimeType(dt) =>
val set = if (nested) "setSafe" else "set"
OutputEmit("", s"$targetVec.$set($idx, $source);")
Expand Down
Loading
Loading