diff --git a/docs/source/user-guide/latest/expressions.md b/docs/source/user-guide/latest/expressions.md index 208a5f3f124..dffaaa10b91 100644 --- a/docs/source/user-guide/latest/expressions.md +++ b/docs/source/user-guide/latest/expressions.md @@ -414,7 +414,7 @@ The type-name conversion functions (`bigint`, `binary`, `boolean`, `date`, `deci | Function | Status | Implementation | Notes | | --- | --- | --- | --- | | `%` | ✅ | Native | | -| `*` | ✅ | Native | DayTime interval multiplication routes through the JVM codegen dispatcher; YearMonth and Calendar interval multiplication fall back | +| `*` | ✅ | Native | YearMonth and DayTime interval multiplication routes through the JVM codegen dispatcher; Calendar interval multiplication falls back | | `+` | ✅ | Native | | | `-` | ✅ | Native | | | `/` | ✅ | Native | | diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 5bd4bfebb77..33b1aed84b5 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -52,7 +52,7 @@ use crate::execution::{ use crate::jvm_bridge::{jni_call, JVMClasses, ShufflePartitionPusher}; use arrow::compute::CastOptions; use arrow::datatypes::{ - DataType, Field, FieldRef, Fields, Schema, TimeUnit, DECIMAL128_MAX_PRECISION, + DataType, Field, FieldRef, Fields, IntervalUnit, Schema, TimeUnit, DECIMAL128_MAX_PRECISION, }; use arrow::ffi_stream::FFI_ArrowArrayStream; use datafusion::functions_aggregate::bit_and_or_xor::{bit_and_udaf, bit_or_udaf, bit_xor_udaf}; @@ -119,8 +119,8 @@ use datafusion::physical_expr::LexOrdering; use crate::parquet::parquet_exec::init_datasource_exec; use arrow::array::{ new_empty_array, Array, ArrayRef, BinaryBuilder, BooleanArray, Date32Array, Decimal128Array, - Float32Array, Float64Array, Int16Array, Int32Array, Int64Array, Int8Array, ListArray, - NullArray, StringBuilder, TimestampMicrosecondArray, + Float32Array, Float64Array, Int16Array, Int32Array, Int64Array, Int8Array, + IntervalYearMonthArray, ListArray, NullArray, StringBuilder, TimestampMicrosecondArray, }; use arrow::buffer::{BooleanBuffer, NullBuffer, OffsetBuffer}; use arrow::row::{OwnedRow, RowConverter, SortField}; @@ -556,6 +556,9 @@ impl PhysicalPlanner { DataType::Time64(TimeUnit::Nanosecond) => { ScalarValue::Time64Nanosecond(None) } + DataType::Interval(IntervalUnit::YearMonth) => { + ScalarValue::IntervalYearMonth(None) + } DataType::Duration(TimeUnit::Microsecond) => { ScalarValue::DurationMicrosecond(None) } @@ -571,9 +574,12 @@ impl PhysicalPlanner { Value::IntVal(value) => match data_type { DataType::Int32 => ScalarValue::Int32(Some(*value)), DataType::Date32 => ScalarValue::Date32(Some(*value)), + DataType::Interval(IntervalUnit::YearMonth) => { + ScalarValue::IntervalYearMonth(Some(*value)) + } dt => { return Err(GeneralError(format!( - "Expected either 'Int32' or 'Date32' for IntVal, but found {dt:?}" + "Expected either 'Int32', 'Date32', or 'Interval(YearMonth)' for IntVal, but found {dt:?}" ))) } }, @@ -4813,6 +4819,10 @@ fn literal_to_array_ref( list_literal.int_values.into(), Some(nulls.clone().into()), ))), + DataType::Interval(IntervalUnit::YearMonth) => Ok(Arc::new(IntervalYearMonthArray::new( + list_literal.int_values.into(), + Some(nulls.clone().into()), + ))), DataType::Timestamp(TimeUnit::Microsecond, None) => { Ok(Arc::new(TimestampMicrosecondArray::new( list_literal.long_values.into(), 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 2ed7e33c904..d8c6ef07d04 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -58,6 +58,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { classOf[TinyIntVector], classOf[SmallIntVector], classOf[IntVector], + classOf[IntervalYearVector], classOf[BigIntVector], classOf[Float4Vector], classOf[Float8Vector], diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index be4bc9c3412..a7b4f6add71 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -305,6 +305,7 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[MakeDTInterval] -> CometMakeDTInterval, classOf[MakeInterval] -> CometMakeInterval, classOf[MultiplyDTInterval] -> CometMultiplyDTInterval, + classOf[MultiplyYMInterval] -> CometMultiplyYMInterval, classOf[TimestampAdd] -> CometTimestampAdd, classOf[TimestampDiff] -> CometTimestampDiff, classOf[MicrosToTimestamp] -> CometMicrosToTimestamp, 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 8f2b9ab6719..b27337db973 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, 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, MultiplyYMInterval, 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 @@ -956,6 +956,8 @@ object CometGetTimestamp extends CometCodegenDispatch[GetTimestamp] object CometMakeYMInterval extends CometCodegenDispatch[MakeYMInterval] +object CometMultiplyYMInterval extends CometCodegenDispatch[MultiplyYMInterval] + object CometMakeDTInterval extends CometCodegenDispatch[MakeDTInterval] object CometMakeInterval extends CometExpressionSerde[MakeInterval] with CodegenDispatchFallback { diff --git a/spark/src/main/scala/org/apache/comet/serde/literals.scala b/spark/src/main/scala/org/apache/comet/serde/literals.scala index b368f80320c..a46c431a48c 100644 --- a/spark/src/main/scala/org/apache/comet/serde/literals.scala +++ b/spark/src/main/scala/org/apache/comet/serde/literals.scala @@ -25,7 +25,7 @@ import org.apache.spark.internal.Logging import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, CreateArray, CreateMap, Expression, KnownNullable, Literal, MapFromArrays} import org.apache.spark.sql.catalyst.util.{ArrayData, MapData, TypeUtils} -import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, CalendarIntervalType, DataType, DateType, DayTimeIntervalType, Decimal, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, NullType, ShortType, StringType, StructType, TimestampNTZType, TimestampType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, CalendarIntervalType, DataType, DateType, DayTimeIntervalType, Decimal, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, NullType, ShortType, StringType, StructType, TimestampNTZType, TimestampType, YearMonthIntervalType} import org.apache.spark.unsafe.types.{CalendarInterval, UTF8String} import com.google.protobuf.ByteString @@ -58,7 +58,11 @@ object CometLiteral extends CometExpressionSerde[Literal] with CometTypeShim wit Unsupported(Some(s"Unsupported literal value for data type ${expr.dataType}")) } else { expr.dataType match { - case _: DayTimeIntervalType => Compatible(None) + // ANSI interval types are deliberately kept out of QueryPlanSerde.supportedDataType so + // they are not claimed as flowing through arbitrary native operators; literals are + // supported here. + case _: DayTimeIntervalType | _: YearMonthIntervalType => + Compatible(None) case dt => Unsupported(Some(s"Unsupported data type $dt")) } } @@ -88,7 +92,8 @@ object CometLiteral extends CometExpressionSerde[Literal] with CometTypeShim wit case _: BooleanType => exprBuilder.setBoolVal(value.asInstanceOf[Boolean]) case _: ByteType => exprBuilder.setByteVal(value.asInstanceOf[Byte]) case _: ShortType => exprBuilder.setShortVal(value.asInstanceOf[Short]) - case _: IntegerType | _: DateType => exprBuilder.setIntVal(value.asInstanceOf[Int]) + case _: IntegerType | _: DateType | _: YearMonthIntervalType => + exprBuilder.setIntVal(value.asInstanceOf[Int]) case _: LongType | _: TimestampType | _: TimestampNTZType | _: DayTimeIntervalType => exprBuilder.setLongVal(value.asInstanceOf[Long]) case dt if isTimeType(dt) => @@ -165,7 +170,7 @@ object CometLiteral extends CometExpressionSerde[Literal] with CometTypeShim wit else null.asInstanceOf[Integer]) listLiteralBuilder.addNullMask(casted != null) }) - case IntegerType | DateType => + case IntegerType | DateType | _: YearMonthIntervalType => array.foreach(v => { val casted = v.asInstanceOf[Integer] listLiteralBuilder.addIntValues(casted) @@ -240,6 +245,9 @@ object CometLiteral extends CometExpressionSerde[Literal] with CometTypeShim wit TimestampType | TimestampNTZType | FloatType | DoubleType | StringType | BinaryType => true case _: DecimalType => true + // Matched as a type rather than a stable identifier: the start/end fields participate in + // `equals`, and every (start, end) pair is carried as the same month count. + case _: YearMonthIntervalType => true case ArrayType(elementType, _) => listLiteralElementSupported(elementType) case _ => false } diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/multiply_ym_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/multiply_ym_interval.sql new file mode 100644 index 00000000000..c64a1a786bb --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/datetime/multiply_ym_interval.sql @@ -0,0 +1,64 @@ +-- 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. + +-- Routes multiply_ym_interval through the codegen dispatcher; produces YearMonthIntervalType. +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=true + +statement +CREATE TABLE test_multiply_ym_interval(y int, m int, i int, l long, f float, d double, dec decimal(10,2)) USING parquet + +statement +INSERT INTO test_multiply_ym_interval VALUES + (1, 2, 2, CAST(3 AS BIGINT), CAST(1.5 AS FLOAT), CAST(2.5 AS DOUBLE), CAST(2.50 AS DECIMAL(10, 2))), + (-1, 1, -2, CAST(-3 AS BIGINT), CAST(-1.5 AS FLOAT), CAST(-2.5 AS DOUBLE), CAST(-2.50 AS DECIMAL(10, 2))), + (0, 6, 0, CAST(0 AS BIGINT), CAST(0.5 AS FLOAT), CAST(0.5 AS DOUBLE), CAST(0.50 AS DECIMAL(10, 2))), + (2, -6, NULL, NULL, NULL, NULL, NULL) + +query +SELECT + make_ym_interval(y, m) * i, + make_ym_interval(y, m) * l, + make_ym_interval(y, m) * f, + make_ym_interval(y, m) * d, + make_ym_interval(y, m) * dec +FROM test_multiply_ym_interval + +-- literal interval input +query +SELECT INTERVAL '1-2' YEAR TO MONTH * i FROM test_multiply_ym_interval + +-- numeric on the left is normalized by Spark to multiply_ym_interval. +query +SELECT i * make_ym_interval(y, m), 2 * INTERVAL '1-2' YEAR TO MONTH +FROM test_multiply_ym_interval + +-- literal multipliers, including half-up rounding for fractional months. +query +SELECT + make_ym_interval(1, 2) * 2, + make_ym_interval(1, 2) * 1.5D, + make_ym_interval(1, 2) * CAST(1.50 AS DECIMAL(10, 2)), + make_ym_interval(-1, 1) * 1.5D + +-- null interval input +query +SELECT make_ym_interval(NULL, m) * 2 FROM test_multiply_ym_interval + +-- 178956970 years and 7 months is Int.MaxValue months. Multiplication overflows regardless of +-- ANSI mode; the lowercase pattern matches Spark 3.x and 4.x error messages. +query expect_error(overflow) +SELECT make_ym_interval(178956970, 7) * 2 diff --git a/spark/src/test/resources/sql-tests/expressions/math/abs.sql b/spark/src/test/resources/sql-tests/expressions/math/abs.sql index ff28bfd3e9e..c32cadd69ed 100644 --- a/spark/src/test/resources/sql-tests/expressions/math/abs.sql +++ b/spark/src/test/resources/sql-tests/expressions/math/abs.sql @@ -56,10 +56,9 @@ INSERT INTO test_abs_iv_overflow VALUES (1, 2, 3, 4.5), (-106751991, -4, 0, -54. query expect_error(overflow) SELECT abs(make_dt_interval(d, h, m, s)) FROM test_abs_iv_overflow --- pinned fallback: NullPropagation folds the ym null into a bare typed literal that CometLiteral --- does not admit (it special-cases only DayTimeIntervalType, literals.scala:63), so the whole --- projection falls back. Pre-existing gap (#5061); this flips to a failure when it is fixed. -query expect_fallback(Unsupported data type YearMonthIntervalType) +-- NullPropagation folds the ym null into a bare typed literal. CometLiteral admits +-- YearMonthIntervalType literals (one of the gaps tracked in #5061), so it stays native. +query SELECT abs(CAST(NULL AS INTERVAL YEAR TO MONTH)) query diff --git a/spark/src/test/scala/org/apache/comet/serde/CometLiteralSuite.scala b/spark/src/test/scala/org/apache/comet/serde/CometLiteralSuite.scala index 746cc8e4b88..c58e65067f5 100644 --- a/spark/src/test/scala/org/apache/comet/serde/CometLiteralSuite.scala +++ b/spark/src/test/scala/org/apache/comet/serde/CometLiteralSuite.scala @@ -123,4 +123,17 @@ class CometLiteralSuite extends CometTestBase with CometTypeShim { assert(!CometLiteral.listLiteralElementSupported(collated)) assert(!encoderAcceptsElement(Array.empty, ArrayType(collated))) } + + // A year-month interval rides in `int_values` as the month count `IntegerType` uses, and + // `literal_to_array_ref` reads those ints back as an `IntervalYearMonthArray`, so the whole + // array literal is serialized directly instead of sending the projection back to Spark. + // `needsExpansion` has no arm for it either, so declining it here is a fallback, not a rewrite. + test("a year-month interval array literal is serialized rather than declined") { + withParquetTable(Seq((1, 2), (3, 4)), "tbl") { + checkSparkAnswerAndOperator( + sql("SELECT array(INTERVAL '1-2' YEAR TO MONTH, INTERVAL '-3' MONTH, NULL) FROM tbl")) + checkSparkAnswerAndOperator( + sql("SELECT array(array(INTERVAL '2-1' YEAR TO MONTH), array()) FROM tbl")) + } + } }