From 72997e9e8f6d234d6811174e174f9cf96c0dd96e Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 7 Aug 2026 22:13:45 +0800 Subject: [PATCH 1/3] fix: preserve CalendarInterval microseconds --- docs/source/user-guide/latest/expressions.md | 2 +- native/Cargo.lock | 1 - native/core/src/execution/serde.rs | 5 +- native/spark-expr/Cargo.toml | 3 +- .../src/datetime_funcs/make_interval.rs | 181 ++++++++++++++---- native/spark-expr/src/datetime_funcs/mod.rs | 2 +- native/spark-expr/src/lib.rs | 8 +- .../comet/vector/CometStructVector.java | 10 + .../CometBatchKernelCodegenInput.scala | 20 +- .../CometBatchKernelCodegenOutput.scala | 19 +- .../org/apache/comet/serde/datetime.scala | 24 +-- .../udf/codegen/CometScalaUDFCodegen.scala | 2 + .../comet/execution/arrow/ArrowWriters.scala | 23 +++ .../apache/spark/sql/comet/util/Utils.scala | 40 +++- .../expressions/datetime/make_interval.sql | 43 ++++- .../datetime/make_interval_ansi.sql | 43 ++++- .../datetime/make_interval_dispatch.sql | 4 +- .../datetime/make_interval_dispatch_ansi.sql | 2 +- .../datetime/try_make_interval.sql | 12 ++ .../org/apache/comet/CometCodegenSuite.scala | 16 +- .../CometDatetimeExpressionBenchmark.scala | 9 +- .../arrow/CometArrowStreamSuite.scala | 17 +- 22 files changed, 367 insertions(+), 119 deletions(-) diff --git a/docs/source/user-guide/latest/expressions.md b/docs/source/user-guide/latest/expressions.md index 192eae807e..9515efaae7 100644 --- a/docs/source/user-guide/latest/expressions.md +++ b/docs/source/user-guide/latest/expressions.md @@ -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 | diff --git a/native/Cargo.lock b/native/Cargo.lock index ae8a2f7be7..eeaa9da5ce 100644 --- a/native/Cargo.lock +++ b/native/Cargo.lock @@ -2066,7 +2066,6 @@ dependencies = [ "datafusion", "datafusion-comet-common", "datafusion-comet-jni-bridge", - "datafusion-spark", "futures", "jni 0.22.4", "num", diff --git a/native/core/src/execution/serde.rs b/native/core/src/execution/serde.rs index f31cb9cd35..31f3553e47 100644 --- a/native/core/src/execution/serde.rs +++ b/native/core/src/execution/serde.rs @@ -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}; @@ -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 diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 6faa9fec4e..a1de8616fe 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -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 } @@ -222,4 +221,4 @@ harness = false [[bench]] name = "cast_int_to_decimal" -harness = false \ No newline at end of file +harness = false diff --git a/native/spark-expr/src/datetime_funcs/make_interval.rs b/native/spark-expr/src/datetime_funcs/make_interval.rs index 9f7b21ca2a..bb721442b6 100644 --- a/native/spark-expr/src/datetime_funcs/make_interval.rs +++ b/native/spark-expr/src/datetime_funcs/make_interval.rs @@ -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 { - self.inner.return_type(arg_types) + fn return_type(&self, _: &[DataType]) -> Result { + Ok(calendar_interval_type()) } fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { - 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::>>()?; + let years = arrays[0].as_any().downcast_ref::().unwrap(); + let months = arrays[1].as_any().downcast_ref::().unwrap(); + let weeks = arrays[2].as_any().downcast_ref::().unwrap(); + let days = arrays[3].as_any().downcast_ref::().unwrap(); + let hours = arrays[4].as_any().downcast_ref::().unwrap(); + let minutes = arrays[5].as_any().downcast_ref::().unwrap(); + let seconds = arrays[6] + .as_any() + .downcast_ref::() + .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 = 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()); } } diff --git a/native/spark-expr/src/datetime_funcs/mod.rs b/native/spark-expr/src/datetime_funcs/mod.rs index 05530f29c2..74f4f4ea12 100644 --- a/native/spark-expr/src/datetime_funcs/mod.rs +++ b/native/spark-expr/src/datetime_funcs/mod.rs @@ -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; diff --git a/native/spark-expr/src/lib.rs b/native/spark-expr/src/lib.rs index 2b5c29befc..20294d90e3 100644 --- a/native/spark-expr/src/lib.rs +++ b/native/spark-expr/src/lib.rs @@ -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::*; diff --git a/spark/src/main/java/org/apache/comet/vector/CometStructVector.java b/spark/src/main/java/org/apache/comet/vector/CometStructVector.java index 259793b831..25514f07f3 100644 --- a/spark/src/main/java/org/apache/comet/vector/CometStructVector.java +++ b/spark/src/main/java/org/apache/comet/vector/CometStructVector.java @@ -27,6 +27,7 @@ import org.apache.arrow.vector.dictionary.DictionaryProvider; import org.apache.arrow.vector.util.TransferPair; import org.apache.spark.sql.vectorized.ColumnVector; +import org.apache.spark.unsafe.types.CalendarInterval; /** * A {@link CometDecodedVector} for Spark struct columns, wrapping an Arrow {@link StructVector} and @@ -58,6 +59,15 @@ public ColumnVector getChild(int i) { return children.get(i); } + @Override + public CalendarInterval getInterval(int rowId) { + if (isNullAt(rowId)) return null; + return new CalendarInterval( + children.get(0).getInt(rowId), + children.get(1).getInt(rowId), + children.get(2).getLong(rowId)); + } + @Override public CometVector slice(int offset, int length) { TransferPair tp = this.valueVector.getTransferPair(this.valueVector.getAllocator()); diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala index 2ed7e33c90..69586898da 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -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, @@ -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 = { @@ -112,6 +113,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { // CometPlainVector exposes `isNullAt`; Arrow-typed fields expose `isNull`. Same semantics. val method = spec.vectorClass match { case cls if wrapsInCometPlainVector(cls) => "isNullAt" + case cls if cls == classOf[StructVector] => "isNullAt" case _ => "isNull" } s" case $ord: return this.col$ord.$method(this.rowIdx);" @@ -146,7 +148,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 { @@ -425,6 +428,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" } @@ -436,9 +440,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;" @@ -471,6 +479,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;" } diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala index 33e6c0c035..996b037c40 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -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 @@ -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);") diff --git a/spark/src/main/scala/org/apache/comet/serde/datetime.scala b/spark/src/main/scala/org/apache/comet/serde/datetime.scala index a74a158e97..0c50613298 100644 --- a/spark/src/main/scala/org/apache/comet/serde/datetime.scala +++ b/spark/src/main/scala/org/apache/comet/serde/datetime.scala @@ -21,7 +21,7 @@ package org.apache.comet.serde import java.util.Locale -import org.apache.spark.sql.catalyst.expressions.{AddMonths, Attribute, Cast, ConvertTimezone, DateAdd, DateDiff, DateFormatClass, DateFromUnixDate, DateSub, DayOfMonth, DayOfWeek, DayOfYear, Days, Expression, FromUTCTimestamp, GetDateField, GetTimestamp, Hour, Hours, LastDay, Literal, MakeDate, MakeDTInterval, MakeInterval, MakeTimestamp, MakeYMInterval, MicrosToTimestamp, MillisToTimestamp, Minute, Month, MonthsBetween, MultiplyDTInterval, NextDay, PreciseTimestampConversion, Quarter, Second, SecondsToTimestamp, TimestampAdd, TimestampDiff, ToUnixTimestamp, ToUTCTimestamp, TruncDate, TruncTimestamp, UnixDate, UnixMicros, UnixMillis, UnixSeconds, UnixTimestamp, WeekDay, WeekOfYear, Year} +import org.apache.spark.sql.catalyst.expressions.{AddMonths, Attribute, ConvertTimezone, DateAdd, DateDiff, DateFormatClass, DateFromUnixDate, DateSub, DayOfMonth, DayOfWeek, DayOfYear, Days, Expression, FromUTCTimestamp, GetDateField, GetTimestamp, Hour, Hours, LastDay, Literal, MakeDate, MakeDTInterval, MakeInterval, MakeTimestamp, MakeYMInterval, MicrosToTimestamp, MillisToTimestamp, Minute, Month, MonthsBetween, MultiplyDTInterval, NextDay, PreciseTimestampConversion, Quarter, Second, SecondsToTimestamp, TimestampAdd, TimestampDiff, ToUnixTimestamp, ToUTCTimestamp, TruncDate, TruncTimestamp, UnixDate, UnixMicros, UnixMillis, UnixSeconds, UnixTimestamp, WeekDay, WeekOfYear, Year} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{CalendarIntervalType, DataType, DateType, DoubleType, FloatType, IntegerType, LongType, StringType, TimestampNTZType, TimestampType} import org.apache.spark.unsafe.types.UTF8String @@ -963,30 +963,14 @@ object CometMakeYMInterval extends CometCodegenDispatch[MakeYMInterval] object CometMakeDTInterval extends CometCodegenDispatch[MakeDTInterval] -object CometMakeInterval extends CometExpressionSerde[MakeInterval] with CodegenDispatchFallback { - private val incompatReason = - "The native implementation converts seconds to `Float64`, which can lose microsecond" + - " precision, and stores time in nanoseconds, which overflows for large time components" + - " (hours, minutes, seconds) that Spark can represent." - - override def getCompatibleNotes(): Seq[String] = Seq( - "Both the default JVM codegen-dispatch path and the native path currently limit the" + - " elapsed-time component to about 292 years in either direction. This only affects" + - " extreme intervals and is tracked in" + - " [#5279](https://github.com/apache/datafusion-comet/issues/5279).") - - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) - - override def getSupportLevel(expr: MakeInterval): SupportLevel = - Incompatible(Some(incompatReason)) +object CometMakeInterval extends CometExpressionSerde[MakeInterval] { + override def getSupportLevel(expr: MakeInterval): SupportLevel = Compatible() override def convert( expr: MakeInterval, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = { - // The explicit return type skips DataFusion's registry coercion, but its kernel needs Float64. - val children = expr.children.updated(6, Cast(expr.secs, DoubleType)) - val childExprs = children.map(exprToProtoInternal(_, inputs, binding)) + val childExprs = expr.children.map(exprToProtoInternal(_, inputs, binding)) val optExpr = scalarFunctionExprToProtoWithReturnType( "make_interval", CalendarIntervalType, diff --git a/spark/src/main/scala/org/apache/comet/udf/codegen/CometScalaUDFCodegen.scala b/spark/src/main/scala/org/apache/comet/udf/codegen/CometScalaUDFCodegen.scala index ef541fb3e2..791ae0a19e 100644 --- a/spark/src/main/scala/org/apache/comet/udf/codegen/CometScalaUDFCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/udf/codegen/CometScalaUDFCodegen.scala @@ -211,6 +211,8 @@ class CometScalaUDFCodegen extends CometUDF with Logging { case list: ListVector => val child = list.getDataVector ArrayColumnSpec(nullable = true, Utils.fromArrowField(child.getField), specFor(child)) + case struct: StructVector if Utils.isCalendarIntervalStructField(struct.getField) => + ScalarColumnSpec(classOf[StructVector], nullable = true) case struct: StructVector => val fieldSpecs = (0 until struct.size()).map { fi => val childVec = struct.getChildByOrdinal(fi).asInstanceOf[ValueVector] diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index 3f47130e5a..5104e92b43 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -25,6 +25,7 @@ import org.apache.arrow.memory.BufferAllocator import org.apache.arrow.vector._ import org.apache.arrow.vector.complex._ import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.GenericInternalRow import org.apache.spark.sql.catalyst.expressions.SpecializedGetters import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.errors.QueryExecutionErrors @@ -87,6 +88,11 @@ private[arrow] object ArrowWriter { case (_: DayTimeIntervalType, vector: DurationVector) => new DurationWriter(vector) case (CalendarIntervalType, vector: IntervalMonthDayNanoVector) => new IntervalMonthDayNanoWriter(vector) + case (CalendarIntervalType, vector: StructVector) => + val children = (0 until vector.size()).map { ordinal => + createFieldWriter(vector.getChildByOrdinal(ordinal)) + } + new CalendarIntervalStructWriter(vector, children.toArray) case (dt, _) => throw QueryExecutionErrors.notSupportTypeError(dt) } @@ -470,6 +476,23 @@ private[arrow] class StructWriter( } } +private[arrow] class CalendarIntervalStructWriter( + valueVector: StructVector, + children: Array[ArrowFieldWriter]) + extends StructWriter(valueVector, children) { + + private val row = new GenericInternalRow(3) + + override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { + valueVector.setIndexDefined(count) + val interval = input.getInterval(ordinal) + row.update(0, interval.months) + row.update(1, interval.days) + row.update(2, interval.microseconds) + children.indices.foreach(i => children(i).write(row, i)) + } +} + private[arrow] class MapWriter( val valueVector: MapVector, val structVector: StructVector, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala index 9d4b0bce88..9ab0ac2f6d 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala @@ -48,6 +48,8 @@ import org.apache.comet.shims.CometTypeShim import org.apache.comet.vector.CometVector object Utils extends CometTypeShim with Logging { + private val calendarIntervalStructKey = "SPARK::calendarInterval::struct" + def getConfPath(confFileName: String): String = { sys.env .get(COMET_CONF_DIR_ENV) @@ -77,6 +79,8 @@ object Utils extends CometTypeShim with Logging { val elementField = field.getChildren().get(0) val elementType = fromArrowField(elementField) ArrayType(elementType, containsNull = elementField.isNullable) + case ArrowType.Struct.INSTANCE if isCalendarIntervalStructField(field) => + CalendarIntervalType case ArrowType.Struct.INSTANCE => val fields = field.getChildren().asScala.map { child => val dt = fromArrowField(child) @@ -165,7 +169,7 @@ object Utils extends CometTypeShim with Logging { // Spark stores DayTimeIntervalType as microseconds in an int64, matching Arrow // Duration(Microsecond) rather than the lossy Interval(DayTime) {days, millis} layout. case _: DayTimeIntervalType => new ArrowType.Duration(TimeUnit.MICROSECOND) - case CalendarIntervalType => new ArrowType.Interval(IntervalUnit.MONTH_DAY_NANO) + case CalendarIntervalType => ArrowType.Struct.INSTANCE case _ => throw new UnsupportedOperationException( s"Unsupported data type: [${dt.getClass.getName}] ${dt.catalogString}") @@ -205,12 +209,46 @@ object Utils extends CometTypeShim with Logging { .add(MapVector.VALUE_NAME, valueType, nullable = valueContainsNull), nullable = false, timeZoneId)).asJava) + case CalendarIntervalType => + val fieldType = new FieldType(nullable, ArrowType.Struct.INSTANCE, null) + val monthsType = new FieldType( + false, + new ArrowType.Int(32, true), + null, + Map(calendarIntervalStructKey -> "true").asJava) + new Field( + name, + fieldType, + Seq( + new Field("months", monthsType, Seq.empty[Field].asJava), + new Field( + "days", + new FieldType(false, new ArrowType.Int(32, true), null), + Seq.empty[Field].asJava), + new Field( + "microseconds", + new FieldType(false, new ArrowType.Int(64, true), null), + Seq.empty[Field].asJava)).asJava) case dataType => val fieldType = new FieldType(nullable, toArrowType(dataType, timeZoneId), null) new Field(name, fieldType, Seq.empty[Field].asJava) } } + def isCalendarIntervalStructField(field: Field): Boolean = { + val children = field.getChildren + def child(index: Int, name: String, bits: Int): Boolean = { + val f = children.get(index) + f.getName == name && f.getType == new ArrowType.Int(bits, true) && !f.isNullable + } + field.getType == ArrowType.Struct.INSTANCE && + children.size == 3 && + child(0, "months", 32) && + child(1, "days", 32) && + child(2, "microseconds", 64) && + children.get(0).getMetadata.getOrDefault(calendarIntervalStructKey, "false") == "true" + } + /** * Maps schema from Spark to Arrow. NOTE: timeZoneId required for TimestampType in StructType */ diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval.sql index 952a9085ee..9226098e7a 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval.sql @@ -15,8 +15,6 @@ -- specific language governing permissions and limitations -- under the License. --- Config: spark.comet.expression.MakeInterval.allowIncompatible=true - statement CREATE TABLE test_make_interval( years int, @@ -29,6 +27,18 @@ CREATE TABLE test_make_interval( statement INSERT INTO test_make_interval VALUES + -- Adapted from Spark's MakeInterval expression tests: + -- https://github.com/apache/spark/blob/v4.2.0/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/IntervalExpressionsSuite.scala#L195-L231 + (0, 0, 0, 0, 0, 0, 0.000000), + (-123, 0, 0, 0, 0, 0, 0.000000), + (0, 0, 123, 0, 0, 0, 0.000000), + (0, 0, 0, 0, 0, 0, -0.123000), + (9999, 11, 0, 31, 23, 59, 59.999999), + (10000, 0, 0, 0, 0, 0, -0.000001), + (-9999, -11, 0, -31, -23, -59, -59.999999), + (-10000, 0, 0, 0, 0, 0, 0.000001), + (0, 0, 0, 0, 2147483647, 2147483647, 2149633277.790647), + (100, 11, 1, 1, 12, 30, 1.001001), (1, 2, 3, 4, 5, 6, 7.123456), (0, 1, 0, 1, 0, 0, 100.000001), (-1, -2, -1, -1, -1, -1, -1.500000), @@ -43,7 +53,28 @@ FROM test_make_interval ORDER BY years query -SELECT make_interval(1, 2), make_interval(3), make_interval() +-- Adapted from Spark's SQL and DataFrame API default-argument tests: +-- https://github.com/apache/spark/blob/v4.2.0/sql/core/src/test/resources/sql-tests/inputs/interval.sql#L81-L90 +-- https://github.com/apache/spark/blob/v4.2.0/sql/core/src/test/scala/org/apache/spark/sql/DateFunctionsSuite.scala#L1284-L1323 +SELECT make_interval(), + make_interval(1), + make_interval(1, 2), + make_interval(1, 2, 3), + make_interval(1, 2, 3, 4), + make_interval(1, 2, 3, 4, 5), + make_interval(1, 2, 3, 4, 5, 6), + make_interval(1, 2, 3, 4, 5, 6, 7.008009) + +query +SELECT make_interval(years), + make_interval(years, months), + make_interval(years, months, weeks), + make_interval(years, months, weeks, days), + make_interval(years, months, weeks, days, hours), + make_interval(years, months, weeks, days, hours, mins), + make_interval(years, months, weeks, days, hours, mins, secs) +FROM test_make_interval +WHERE years = 100 query SELECT make_interval(0, 1, 0, 1, 0, 0, 100.000001) @@ -51,16 +82,16 @@ SELECT make_interval(0, 1, 0, 1, 0, 0, 100.000001) query SELECT make_interval(2147483647) -query ignore(https://github.com/apache/datafusion-comet/issues/5131) +query SELECT make_interval(1, 2, 3, 4, 0, 0, 123456789012.123456) query SELECT make_interval(0, 0, 0, 0, 0, 0, 999999999.999999) -query ignore(https://github.com/apache/datafusion-comet/issues/5131) +query SELECT make_interval(0, 0, 0, 0, 0, 0, 999999999.000001) -query ignore(https://github.com/apache/datafusion-comet/issues/5131) +query SELECT make_interval(0, 0, 0, 0, 2562048) query diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_ansi.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_ansi.sql index 175b9ecd08..f3e2f8ca32 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_ansi.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_ansi.sql @@ -15,27 +15,54 @@ -- specific language governing permissions and limitations -- under the License. --- Native ANSI execution must preserve Spark's overflow exception. -- Config: spark.sql.ansi.enabled=true --- Config: spark.comet.expression.MakeInterval.allowIncompatible=true statement -CREATE TABLE test_make_interval_ansi(years int) USING parquet +CREATE TABLE test_make_interval_ansi( + id int, + years int, + months int, + weeks int, + days int, + hours int, + mins int, + secs decimal(18, 6)) USING parquet statement -INSERT INTO test_make_interval_ansi VALUES (NULL) +INSERT INTO test_make_interval_ansi VALUES + -- Adapted from Spark's ANSI MakeInterval expression tests: + -- https://github.com/apache/spark/blob/v4.2.0/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/IntervalExpressionsSuite.scala#L233-L279 + (0, NULL, 0, 0, 0, 0, 0, 0.000000), + (1, 0, 0, 0, 0, 0, 0, 0.000000), + (2, -123, 0, 0, 0, 0, 0, 0.000000), + (3, 0, 0, 123, 0, 0, 0, 0.000000), + (4, 0, 0, 0, 0, 0, 0, -0.123000), + (5, 9999, 11, 0, 31, 23, 59, 59.999999), + (6, 10000, 0, 0, 0, 0, 0, -0.000001), + (7, -9999, -11, 0, -31, -23, -59, -59.999999), + (8, -10000, 0, 0, 0, 0, 0, 0.000001), + (9, 0, 0, 0, 0, 2147483647, 2147483647, 2149633277.790647), + (10, 2147483647, 0, 0, 0, 0, 0, 0.000000), + (11, 0, 0, 2147483647, 0, 0, 0, 0.000000) query SELECT make_interval(1, 2, 3, 4, 5, 6, 7.123456) query -SELECT make_interval(years) FROM test_make_interval_ansi +SELECT make_interval(years, months, weeks, days, hours, mins, secs) +FROM test_make_interval_ansi +WHERE id BETWEEN 0 AND 9 +ORDER BY id query expect_error(overflow. If necessary set) -SELECT make_interval(2147483647) +SELECT make_interval(years) +FROM test_make_interval_ansi +WHERE id = 10 query expect_error(overflow. If necessary set) -SELECT make_interval(0, 0, 2147483647) +SELECT make_interval(0, 0, weeks) +FROM test_make_interval_ansi +WHERE id = 11 -query ignore(https://github.com/apache/datafusion-comet/issues/5131) +query SELECT make_interval(0, 0, 0, 0, 2562048) diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch.sql index ee2d5e8160..57747625f2 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch.sql @@ -15,8 +15,6 @@ -- specific language governing permissions and limitations -- under the License. --- With allowIncompatible unset, MakeInterval uses Spark's JVM codegen dispatcher. - statement CREATE TABLE test_make_interval_dispatch( years int, @@ -43,7 +41,7 @@ FROM test_make_interval_dispatch WHERE hours != 2562048 ORDER BY years -query ignore(https://github.com/apache/datafusion-comet/issues/5279) +query SELECT make_interval(0, 0, 0, 0, hours) FROM test_make_interval_dispatch WHERE hours = 2562048 diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch_ansi.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch_ansi.sql index f199a8fb66..1d9d885cad 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch_ansi.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_interval_dispatch_ansi.sql @@ -42,7 +42,7 @@ SELECT make_interval(0, 0, weeks) FROM test_make_interval_dispatch_ansi WHERE weeks = 2147483647 -query ignore(https://github.com/apache/datafusion-comet/issues/5279) +query SELECT make_interval(0, 0, 0, 0, hours) FROM test_make_interval_dispatch_ansi WHERE hours = 2562048 diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/try_make_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/try_make_interval.sql index 5e481cb9ff..2951a375bf 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/try_make_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/try_make_interval.sql @@ -31,9 +31,21 @@ CREATE TABLE test_try_make_interval( statement INSERT INTO test_try_make_interval VALUES (1, 2, 3, 4, 5, 6, 7.123456), + (0, 0, 0, 0, 2562048, 0, 999999999.000001), (2147483647, 0, 0, 0, 0, 0, 0.000000) query SELECT try_make_interval(years, months, weeks, days, hours, mins, secs) FROM test_try_make_interval ORDER BY years + +query +-- Adapted from Spark's try_make_interval default-argument API tests: +-- https://github.com/apache/spark/blob/v4.2.0/sql/connect/client/jvm/src/test/scala/org/apache/spark/sql/PlanGenerationTestSuite.scala#L2037-L2076 +SELECT try_make_interval(1), + try_make_interval(1, 2), + try_make_interval(1, 2, 3), + try_make_interval(1, 2, 3, 4), + try_make_interval(1, 2, 3, 4, 5), + try_make_interval(1, 2, 3, 4, 5, 6), + try_make_interval(1, 2, 3, 4, 5, 6, 7.008009) diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 5806cb3501..bc108faf19 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -26,6 +26,7 @@ import org.apache.spark.{SparkConf, SparkEnv, TaskContext} import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.api.java.UDF1 import org.apache.spark.sql.catalyst.expressions.{BoundReference, CreateArray, CreateMap, CreateNamedStruct, Expression, Literal, MapConcat} +import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ @@ -82,18 +83,25 @@ class CometCodegenSuite } test("codegen kernel round-trips CalendarIntervalType") { - val input = new IntervalMonthDayNanoVector("in", CometArrowAllocator) + val input = Utils + .toArrowField("in", CalendarIntervalType, nullable = true, "UTC") + .createVector(CometArrowAllocator) + .asInstanceOf[org.apache.arrow.vector.complex.StructVector] val field = CometBatchKernelCodegen.toFfiArrowField("out", CalendarIntervalType, nullable = true) val output = CometBatchKernelCodegen.allocateOutput(field, 2, 0) try { input.allocateNew() - input.setSafe(0, 14, -3, 1234567000L) + input.setIndexDefined(0) + input.getChild("months").asInstanceOf[IntVector].setSafe(0, 14) + input.getChild("days").asInstanceOf[IntVector].setSafe(0, -3) + input.getChild("microseconds").asInstanceOf[BigIntVector].setSafe(0, Long.MaxValue) input.setNull(1) input.setValueCount(2) val expr = BoundReference(0, CalendarIntervalType, nullable = true) - val spec = ArrowColumnSpec(classOf[IntervalMonthDayNanoVector], nullable = true) + val spec = + ArrowColumnSpec(classOf[org.apache.arrow.vector.complex.StructVector], nullable = true) val kernel = CometBatchKernelCodegen.compile(expr, IndexedSeq(spec)).newInstance() kernel.init(0) kernel.process(Array(input), output, 2) @@ -103,7 +111,7 @@ class CometCodegenSuite val actual = comet.getInterval(0) assert(actual.months === 14) assert(actual.days === -3) - assert(actual.microseconds === 1234567L) + assert(actual.microseconds === Long.MaxValue) assert(comet.getInterval(1) == null) } finally { output.close() diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometDatetimeExpressionBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometDatetimeExpressionBenchmark.scala index a9d5b35e17..ec558f9a7b 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometDatetimeExpressionBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometDatetimeExpressionBenchmark.scala @@ -165,18 +165,11 @@ object CometDatetimeExpressionBenchmark extends CometBenchmarkBase { consumeIntervals() } } - benchmark.addCase("Comet (codegen dispatch)") { _ => + benchmark.addCase("Comet") { _ => withSQLConf(cometConfigs.toSeq: _*) { consumeIntervals() } } - benchmark.addCase("Comet (native)") { _ => - val configs = - cometConfigs ++ Map(CometConf.getExprAllowIncompatConfigKey("MakeInterval") -> "true") - withSQLConf(configs.toSeq: _*) { - consumeIntervals() - } - } benchmark.run() } } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala index 50f723d0bf..56938c939d 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala @@ -25,7 +25,8 @@ import org.scalatest.funsuite.AnyFunSuite import org.scalatest.matchers.should.Matchers import org.apache.arrow.memory.RootAllocator -import org.apache.arrow.vector.{BigIntVector, IntervalMonthDayNanoVector, IntVector, VectorSchemaRoot} +import org.apache.arrow.vector.{BigIntVector, IntVector, VectorSchemaRoot} +import org.apache.arrow.vector.complex.StructVector import org.apache.arrow.vector.types.pojo.{ArrowType, Field, FieldType, Schema} import org.apache.spark.sql.catalyst.expressions.GenericInternalRow import org.apache.spark.sql.comet.util.Utils @@ -61,19 +62,19 @@ class CometArrowStreamSuite extends AnyFunSuite with Matchers { Utils.fromArrowField(field) shouldBe CalendarIntervalType val root = VectorSchemaRoot.create(new Schema(Seq(field).asJava), allocator) try { - val expected = new CalendarInterval(14, -3, 1234567L) + val expected = new CalendarInterval(14, -3, Long.MaxValue) val writer = ArrowWriter.create(root) writer.write(new GenericInternalRow(Array[Any](expected))) writer.write(new GenericInternalRow(Array[Any](null))) writer.finish() - val arrow = root.getVector(0).asInstanceOf[IntervalMonthDayNanoVector] - IntervalMonthDayNanoVector.getMonths(arrow.getDataBuffer, 0) shouldBe expected.months - IntervalMonthDayNanoVector.getDays(arrow.getDataBuffer, 0) shouldBe expected.days - IntervalMonthDayNanoVector.getNanoseconds(arrow.getDataBuffer, 0) shouldBe - expected.microseconds * 1000L + val arrow = root.getVector(0).asInstanceOf[StructVector] + arrow.getChild("months").asInstanceOf[IntVector].get(0) shouldBe expected.months + arrow.getChild("days").asInstanceOf[IntVector].get(0) shouldBe expected.days + arrow.getChild("microseconds").asInstanceOf[BigIntVector].get(0) shouldBe + expected.microseconds - val comet = new CometPlainVector(arrow, false) + val comet = CometVector.getVector(arrow, null) comet.getInterval(0) shouldBe expected comet.getInterval(1) shouldBe null } finally { From a0cb4ddbe917596e46549300019323452e7269cf Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 7 Aug 2026 23:14:18 +0800 Subject: [PATCH 2/3] scalafix --- .../org/apache/comet/vector/CometStructVector.java | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/spark/src/main/java/org/apache/comet/vector/CometStructVector.java b/spark/src/main/java/org/apache/comet/vector/CometStructVector.java index 25514f07f3..259793b831 100644 --- a/spark/src/main/java/org/apache/comet/vector/CometStructVector.java +++ b/spark/src/main/java/org/apache/comet/vector/CometStructVector.java @@ -27,7 +27,6 @@ import org.apache.arrow.vector.dictionary.DictionaryProvider; import org.apache.arrow.vector.util.TransferPair; import org.apache.spark.sql.vectorized.ColumnVector; -import org.apache.spark.unsafe.types.CalendarInterval; /** * A {@link CometDecodedVector} for Spark struct columns, wrapping an Arrow {@link StructVector} and @@ -59,15 +58,6 @@ public ColumnVector getChild(int i) { return children.get(i); } - @Override - public CalendarInterval getInterval(int rowId) { - if (isNullAt(rowId)) return null; - return new CalendarInterval( - children.get(0).getInt(rowId), - children.get(1).getInt(rowId), - children.get(2).getLong(rowId)); - } - @Override public CometVector slice(int offset, int length) { TransferPair tp = this.valueVector.getTransferPair(this.valueVector.getAllocator()); From 9b764d83eb28931ddd8a39f5b699baa526727364 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sat, 8 Aug 2026 01:50:26 +0800 Subject: [PATCH 3/3] fix: use Arrow null check for struct vectors --- .../codegen/CometBatchKernelCodegenInput.scala | 7 +------ .../apache/comet/CometCodegenSourceSuite.scala | 18 ++++++++++++++++++ 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala index 69586898da..964a7b1f64 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -110,12 +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 cls if cls == classOf[StructVector] => "isNullAt" - case _ => "isNull" - } + val method = nullCheckMethod(spec) s" case $ord: return this.col$ord.$method(this.rowIdx);" } } diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala index adc4440915..f24a6abbec 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala @@ -114,6 +114,24 @@ class CometCodegenSourceSuite extends AnyFunSuite { s"expected nullable isNullAt to delegate to the Arrow vector; got:\n$src") } + test("nullable struct delegates isNullAt to Arrow isNull") { + val structType = StructType(Seq(StructField("i", IntegerType, nullable = true))) + val structSpec = StructColumnSpec( + nullable = true, + fields = Seq( + StructFieldSpec( + "i", + IntegerType, + nullable = true, + ScalarColumnSpec( + CometBatchKernelCodegen.vectorClassBySimpleName("IntVector"), + nullable = true)))) + val src = gen(BoundReference(0, structType, nullable = true), structSpec) + assert( + src.contains("case 0: return this.col0.isNull(this.rowIdx);"), + s"expected nullable StructVector to use Arrow isNull; got:\n$src") + } + test("VarCharVector getUTF8String uses zero-copy fromAddress") { val expr = Length(BoundReference(0, StringType, nullable = true)) val src = gen(expr, nullableString)