From 9e944735b9b38c3231e897c842e27bd30a3f3bd2 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Mon, 20 Jul 2026 21:33:37 +0800 Subject: [PATCH 1/6] feat: support interval codegen dispatch --- .../CometBatchKernelCodegenInput.scala | 20 ++++++++++++------- .../CometBatchKernelCodegenOutput.scala | 5 +++-- .../CometSpecializedGettersDispatch.scala | 5 +++-- .../comet/serde/operator/CometSink.scala | 18 ++++++++++++----- .../udf/codegen/CometScalaUDFCodegen.scala | 2 +- .../sql/comet/CometLocalTableScanExec.scala | 4 ++-- .../shuffle/CometShuffleExchangeExec.scala | 5 +++-- .../expressions/datetime/make_dt_interval.sql | 16 +++++++++++++++ .../expressions/datetime/make_ym_interval.sql | 16 +++++++++++++++ 9 files changed, 70 insertions(+), 21 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 1d94bdc8dc4..d27f34a5860 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -65,7 +65,9 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { classOf[DurationVector], classOf[TimeNanoVector], classOf[TimeStampMicroVector], - classOf[TimeStampMicroTZVector]) + classOf[TimeStampMicroTZVector], + classOf[IntervalYearVector], + classOf[DurationVector]) private val cometPlainVectorName: String = classOf[CometPlainVector].getName /** Emit kernel typed-vector field declarations for every level of every input column. */ @@ -130,7 +132,8 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { } val intCases = withOrd.collect { case (ArrowColumnSpec(cls, _), ord) - if cls == classOf[IntVector] || cls == classOf[DateDayVector] => + if cls == classOf[IntVector] || cls == classOf[DateDayVector] || + cls == classOf[IntervalYearVector] => s" case $ord: return this.col$ord.getInt(this.rowIdx);" } val longCases = withOrd.collect { @@ -139,7 +142,8 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { cls == classOf[DurationVector] || cls == classOf[TimeNanoVector] || cls == classOf[TimeStampMicroVector] || - cls == classOf[TimeStampMicroTZVector] => + cls == classOf[TimeStampMicroTZVector] || + cls == classOf[DurationVector] => s" case $ord: return this.col$ord.getLong(this.rowIdx);" } val floatCases = withOrd.collect { @@ -595,7 +599,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { case BooleanType => s"getBoolean($idx)" case ByteType => s"getByte($idx)" case ShortType => s"getShort($idx)" - case IntegerType | DateType => s"getInt($idx)" + case IntegerType | DateType | _: YearMonthIntervalType => s"getInt($idx)" case LongType | TimestampType | TimestampNTZType | _: DayTimeIntervalType => s"getLong($idx)" case dt if isTimeType(dt) => s"getLong($idx)" @@ -694,7 +698,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { | public short getShort(int i) { | return $childField.getShort(startIndex + i); | }""".stripMargin - case IntegerType | DateType => + case IntegerType | DateType | _: YearMonthIntervalType => s""" @Override | public int getInt(int i) { | return $childField.getInt(startIndex + i); @@ -855,7 +859,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { s" case $fi: return ${path}_f$fi.getByte(this.rowIdx);" case ShortType => s" case $fi: return ${path}_f$fi.getShort(this.rowIdx);" - case IntegerType | DateType => + case IntegerType | DateType | _: YearMonthIntervalType => s" case $fi: return ${path}_f$fi.getInt(this.rowIdx);" case LongType | TimestampType | TimestampNTZType | _: DayTimeIntervalType => s" case $fi: return ${path}_f$fi.getLong(this.rowIdx);" @@ -905,7 +909,9 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { fieldReadScalar(fi, ShortType, f.nullable) } val intCases = scalarOrd.collect { - case (f, fi) if f.sparkType == IntegerType || f.sparkType == DateType => + case (f, fi) + if f.sparkType == IntegerType || f.sparkType == DateType || + f.sparkType.isInstanceOf[YearMonthIntervalType] => fieldReadScalar(fi, IntegerType, f.nullable) } val longCases = scalarOrd.collect { 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 4160c478d13..00c62805433 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -401,8 +401,9 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { case BooleanType => s"$target.getBoolean($idx)" case ByteType => s"$target.getByte($idx)" case ShortType => s"$target.getShort($idx)" - case IntegerType | DateType => s"$target.getInt($idx)" - case LongType | TimestampType | TimestampNTZType => s"$target.getLong($idx)" + case IntegerType | DateType | _: YearMonthIntervalType => s"$target.getInt($idx)" + case LongType | TimestampType | TimestampNTZType | _: DayTimeIntervalType => + s"$target.getLong($idx)" case dt if isTimeType(dt) => s"$target.getLong($idx)" case FloatType => s"$target.getFloat($idx)" case DoubleType => s"$target.getDouble($idx)" diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometSpecializedGettersDispatch.scala b/spark/src/main/scala/org/apache/comet/codegen/CometSpecializedGettersDispatch.scala index 2f81c58c06e..3eeca4e4041 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometSpecializedGettersDispatch.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometSpecializedGettersDispatch.scala @@ -40,8 +40,9 @@ private[codegen] object CometSpecializedGettersDispatch { case BooleanType => java.lang.Boolean.valueOf(g.getBoolean(ordinal)) case ByteType => java.lang.Byte.valueOf(g.getByte(ordinal)) case ShortType => java.lang.Short.valueOf(g.getShort(ordinal)) - case IntegerType | DateType => java.lang.Integer.valueOf(g.getInt(ordinal)) - case LongType | TimestampType | TimestampNTZType => + case IntegerType | DateType | _: YearMonthIntervalType => + java.lang.Integer.valueOf(g.getInt(ordinal)) + case LongType | TimestampType | TimestampNTZType | _: DayTimeIntervalType => java.lang.Long.valueOf(g.getLong(ordinal)) case FloatType => java.lang.Float.valueOf(g.getFloat(ordinal)) case DoubleType => java.lang.Double.valueOf(g.getDouble(ordinal)) diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala index ed69f82e2ab..f2026013cbf 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala @@ -26,7 +26,7 @@ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec import org.apache.spark.sql.execution.exchange.ReusedExchangeExec -import org.apache.spark.sql.types.DataType +import org.apache.spark.sql.types.{ArrayType, DataType, DayTimeIntervalType, MapType, StructType, YearMonthIntervalType} import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.withFallbackReason @@ -43,6 +43,16 @@ abstract class CometSink[T <: SparkPlan] extends CometOperatorSerde[T] { override def enabledConfig: Option[ConfigEntry[Boolean]] = None + protected final def supportedSinkDataType(dt: DataType): Boolean = dt match { + case _: YearMonthIntervalType | _: DayTimeIntervalType => true + case StructType(fields) => + fields.nonEmpty && fields.forall(f => supportedSinkDataType(f.dataType)) + case ArrayType(elementType, _) => supportedSinkDataType(elementType) + case MapType(keyType, valueType, _) => + supportedSinkDataType(keyType) && supportedSinkDataType(valueType) + case _ => supportedDataType(dt) + } + /** * The data type to declare for a scan output field. Overridden by sinks whose source carries * non-null nested child fields that must be widened to match the planned kernel output types @@ -54,8 +64,7 @@ abstract class CometSink[T <: SparkPlan] extends CometOperatorSerde[T] { op: T, builder: Operator.Builder, childOp: OperatorOuterClass.Operator*): Option[OperatorOuterClass.Operator] = { - val supportedTypes = - op.output.forall(a => supportedDataType(a.dataType, allowComplex = true)) + val supportedTypes = op.output.forall(a => supportedSinkDataType(a.dataType)) if (!supportedTypes) { withFallbackReason(op, "Unsupported data type") @@ -123,8 +132,7 @@ object CometExchangeSink extends CometSink[SparkPlan] { private def convertToShuffleScan( op: SparkPlan, builder: Operator.Builder): Option[OperatorOuterClass.Operator] = { - val supportedTypes = - op.output.forall(a => supportedDataType(a.dataType, allowComplex = true)) + val supportedTypes = op.output.forall(a => supportedSinkDataType(a.dataType)) if (!supportedTypes) { withFallbackReason(op, "Unsupported data type for shuffle direct read") 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 145698102ed..9127bff5ac1 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 @@ -220,7 +220,7 @@ class CometScalaUDFCodegen extends CometUDF with Logging { case _: BitVector | _: TinyIntVector | _: SmallIntVector | _: IntVector | _: BigIntVector | _: Float4Vector | _: Float8Vector | _: DecimalVector | _: VarCharVector | _: VarBinaryVector | _: DateDayVector | _: DurationVector | _: TimeStampMicroVector | - _: TimeStampMicroTZVector => + _: TimeStampMicroTZVector | _: IntervalYearVector => ScalarColumnSpec(v.getClass.asInstanceOf[Class[_ <: ValueVector]], nullable = true) case other => throw new UnsupportedOperationException( diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala index 28ff6c85c2a..d5bd16bdae6 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala @@ -31,7 +31,7 @@ import org.apache.spark.sql.comet.execution.arrow.{CometArrowStream, CometNative import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.{LeafExecNode, LocalTableScanExec} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} -import org.apache.spark.sql.types.{DataType, NullType, StructType} +import org.apache.spark.sql.types.{DataType, DayTimeIntervalType, NullType, StructType, YearMonthIntervalType} import com.google.common.base.Objects @@ -138,7 +138,7 @@ object CometLocalTableScanExec extends CometSink[LocalTableScanExec] with DataTy dt: DataType, name: String, fallbackReasons: ListBuffer[String]): Boolean = dt match { - case _: NullType => true + case _: NullType | _: YearMonthIntervalType | _: DayTimeIntervalType => true case _ => super.isTypeSupported(dt, name, fallbackReasons) } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index c9fe324bd81..71d17846305 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -41,7 +41,7 @@ import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeExec, ShuffleExchangeLike, ShuffleOrigin} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics, SQLShuffleReadMetricsReporter, SQLShuffleWriteMetricsReporter} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, NullType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, DataType, DateType, DayTimeIntervalType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, NullType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType, YearMonthIntervalType} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.spark.util.MutablePair import org.apache.spark.util.collection.unsafe.sort.{PrefixComparators, RecordComparator} @@ -412,7 +412,8 @@ object CometShuffleExchangeExec def supportedSerializableDataType(dt: DataType): Boolean = dt match { case _: BooleanType | _: ByteType | _: ShortType | _: IntegerType | _: LongType | _: FloatType | _: DoubleType | _: StringType | _: BinaryType | _: TimestampType | - _: TimestampNTZType | _: DecimalType | _: DateType | _: NullType => + _: TimestampNTZType | _: DecimalType | _: DateType | _: NullType | + _: YearMonthIntervalType | _: DayTimeIntervalType => true case dt if isTimeType(dt) => true diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql index 2c61040de93..38e1757dfcc 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql @@ -15,6 +15,9 @@ -- specific language governing permissions and limitations -- under the License. +-- Config: spark.comet.exec.localTableScan.enabled=true +-- Config: spark.comet.exec.shuffle.mode=native + -- Routes make_dt_interval through the codegen dispatcher; produces DayTimeIntervalType. statement @@ -34,6 +37,19 @@ SELECT make_dt_interval(1, 2, 3, 4.5), make_dt_interval(0, 0, 0, 0) query SELECT make_dt_interval(1), make_dt_interval(1, 2), make_dt_interval() +-- nested interval output through LocalTableScan and codegen dispatch +query +SELECT transform(a, x -> x) AS result +FROM VALUES + (array(make_dt_interval(1, 2, 3, 4.5), CAST(NULL AS INTERVAL DAY TO SECOND))) +AS t(a) + +-- interval output through native shuffle +query +SELECT d, h, mi, s, make_dt_interval(d, h, mi, s) AS i +FROM test_mdi +DISTRIBUTE BY d + -- overflow: days * MICROS_PER_DAY exceeds the int64 microsecond range. makeDayTimeInterval throws -- unconditionally (not ANSI-gated); this confirms the dispatched codegen path propagates Spark's -- exception. The pattern is the lowercase word so it matches every version: Spark 4.x raises diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql index a642b4925db..4d557b07d92 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql @@ -15,6 +15,9 @@ -- specific language governing permissions and limitations -- under the License. +-- Config: spark.comet.exec.localTableScan.enabled=true +-- Config: spark.comet.exec.shuffle.mode=native + -- Routes make_ym_interval through the codegen dispatcher; produces YearMonthIntervalType. statement @@ -34,6 +37,19 @@ SELECT make_ym_interval(1, 2), make_ym_interval(0, 0), make_ym_interval(-5, 11) query SELECT make_ym_interval(3), make_ym_interval() +-- nested interval output through LocalTableScan and codegen dispatch +query +SELECT transform(a, x -> x) AS result +FROM VALUES + (array(make_ym_interval(1, 2), CAST(NULL AS INTERVAL YEAR TO MONTH))) +AS t(a) + +-- interval output through native shuffle +query +SELECT y, m, make_ym_interval(y, m) AS i +FROM test_myi +DISTRIBUTE BY y + -- overflow: years * 12 exceeds Int range. makeYearMonthInterval throws unconditionally (not -- ANSI-gated); this confirms the dispatched codegen path propagates Spark's exception. The -- pattern is the lowercase word so it matches every version: Spark 4.x raises From 72deb3c20c9544ea7babb8d2cdd459c0d66d4ca7 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 24 Jul 2026 23:48:41 +0800 Subject: [PATCH 2/6] add struct and map with nested interval type --- .../sql-tests/expressions/datetime/make_dt_interval.sql | 7 +++++-- .../sql-tests/expressions/datetime/make_ym_interval.sql | 7 +++++-- 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql index 38e1757dfcc..c8fab97a80c 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_dt_interval.sql @@ -44,9 +44,12 @@ FROM VALUES (array(make_dt_interval(1, 2, 3, 4.5), CAST(NULL AS INTERVAL DAY TO SECOND))) AS t(a) --- interval output through native shuffle +-- top-level, struct, and map interval output through native shuffle query -SELECT d, h, mi, s, make_dt_interval(d, h, mi, s) AS i +SELECT d, h, mi, s, + make_dt_interval(d, h, mi, s) AS i, + named_struct('i', make_dt_interval(d, h, mi, s)) AS st, + map('i', make_dt_interval(d, h, mi, s)) AS m FROM test_mdi DISTRIBUTE BY d diff --git a/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql b/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql index 4d557b07d92..b60e4efc521 100644 --- a/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql +++ b/spark/src/test/resources/sql-tests/expressions/datetime/make_ym_interval.sql @@ -44,9 +44,12 @@ FROM VALUES (array(make_ym_interval(1, 2), CAST(NULL AS INTERVAL YEAR TO MONTH))) AS t(a) --- interval output through native shuffle +-- top-level, struct, and map interval output through native shuffle query -SELECT y, m, make_ym_interval(y, m) AS i +SELECT y, m, + make_ym_interval(y, m) AS i, + named_struct('i', make_ym_interval(y, m)) AS s, + map('i', make_ym_interval(y, m)) AS m FROM test_myi DISTRIBUTE BY y From 5bc219280662a9eaa0314ea4d6aaf1ecbde519df Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sat, 25 Jul 2026 00:57:53 +0800 Subject: [PATCH 3/6] Centralize data type support predicates --- .../apache/comet/serde/QueryPlanSerde.scala | 47 +++++++++----- .../comet/serde/operator/CometSink.scala | 18 ++---- .../sql/comet/CometLocalTableScanExec.scala | 20 +++--- .../shuffle/CometShuffleExchangeExec.scala | 63 +++---------------- .../apache/comet/CometExpressionSuite.scala | 28 +++++++++ 5 files changed, 88 insertions(+), 88 deletions(-) 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 19ac36b0405..2d23147709f 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -504,21 +504,38 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { } } - def supportedDataType(dt: DataType, allowComplex: Boolean = false): Boolean = dt match { - case _: ByteType | _: ShortType | _: IntegerType | _: LongType | _: FloatType | - _: DoubleType | _: StringType | _: BinaryType | _: TimestampType | _: TimestampNTZType | - _: DecimalType | _: DateType | _: BooleanType | _: NullType => - true - case dt if isTimeType(dt) => - true - case s: StructType if allowComplex => - s.fields.nonEmpty && s.fields.map(_.dataType).forall(supportedDataType(_, allowComplex)) - case a: ArrayType if allowComplex => - supportedDataType(a.elementType, allowComplex) - case m: MapType if allowComplex => - supportedDataType(m.keyType, allowComplex) && supportedDataType(m.valueType, allowComplex) - case _ => - false + def supportedDataType( + dt: DataType, + allowComplex: Boolean = false, + allowIntervals: Boolean = false, + allowTimeType: Boolean = true, + allowAnyStringType: Boolean = true, + allowDuplicateStructFieldNames: Boolean = true): Boolean = { + def supported(dt: DataType): Boolean = dt match { + case _: ByteType | _: ShortType | _: IntegerType | _: LongType | _: FloatType | + _: DoubleType | _: BinaryType | _: TimestampType | _: TimestampNTZType | + _: DecimalType | _: DateType | _: BooleanType | _: NullType => + true + case st: StringType if allowAnyStringType || st == StringType => + true + case _: YearMonthIntervalType | _: DayTimeIntervalType if allowIntervals => + true + case dt if allowTimeType && isTimeType(dt) => + true + case s: StructType if allowComplex => + s.fields.nonEmpty && + (allowDuplicateStructFieldNames || + s.fields.map(_.name).distinct.length == s.fields.length) && + s.fields.forall(f => supported(f.dataType)) + case a: ArrayType if allowComplex => + supported(a.elementType) + case m: MapType if allowComplex => + supported(m.keyType) && supported(m.valueType) + case _ => + false + } + + supported(dt) } /** diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala index f2026013cbf..c5fc0e4858a 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometSink.scala @@ -26,7 +26,7 @@ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec import org.apache.spark.sql.execution.exchange.ReusedExchangeExec -import org.apache.spark.sql.types.{ArrayType, DataType, DayTimeIntervalType, MapType, StructType, YearMonthIntervalType} +import org.apache.spark.sql.types.DataType import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.withFallbackReason @@ -43,16 +43,6 @@ abstract class CometSink[T <: SparkPlan] extends CometOperatorSerde[T] { override def enabledConfig: Option[ConfigEntry[Boolean]] = None - protected final def supportedSinkDataType(dt: DataType): Boolean = dt match { - case _: YearMonthIntervalType | _: DayTimeIntervalType => true - case StructType(fields) => - fields.nonEmpty && fields.forall(f => supportedSinkDataType(f.dataType)) - case ArrayType(elementType, _) => supportedSinkDataType(elementType) - case MapType(keyType, valueType, _) => - supportedSinkDataType(keyType) && supportedSinkDataType(valueType) - case _ => supportedDataType(dt) - } - /** * The data type to declare for a scan output field. Overridden by sinks whose source carries * non-null nested child fields that must be widened to match the planned kernel output types @@ -64,7 +54,8 @@ abstract class CometSink[T <: SparkPlan] extends CometOperatorSerde[T] { op: T, builder: Operator.Builder, childOp: OperatorOuterClass.Operator*): Option[OperatorOuterClass.Operator] = { - val supportedTypes = op.output.forall(a => supportedSinkDataType(a.dataType)) + val supportedTypes = op.output.forall(a => + supportedDataType(a.dataType, allowComplex = true, allowIntervals = true)) if (!supportedTypes) { withFallbackReason(op, "Unsupported data type") @@ -132,7 +123,8 @@ object CometExchangeSink extends CometSink[SparkPlan] { private def convertToShuffleScan( op: SparkPlan, builder: Operator.Builder): Option[OperatorOuterClass.Operator] = { - val supportedTypes = op.output.forall(a => supportedSinkDataType(a.dataType)) + val supportedTypes = op.output.forall(a => + supportedDataType(a.dataType, allowComplex = true, allowIntervals = true)) if (!supportedTypes) { withFallbackReason(op, "Unsupported data type for shuffle direct read") diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala index d5bd16bdae6..9495a5b96d9 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala @@ -31,13 +31,14 @@ import org.apache.spark.sql.comet.execution.arrow.{CometArrowStream, CometNative import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.{LeafExecNode, LocalTableScanExec} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} -import org.apache.spark.sql.types.{DataType, DayTimeIntervalType, NullType, StructType, YearMonthIntervalType} +import org.apache.spark.sql.types.{DataType, StructType} import com.google.common.base.Objects import org.apache.comet.{CometConf, ConfigEntry, DataTypeSupport} import org.apache.comet.CometSparkSessionExtensions.withFallbackReason import org.apache.comet.serde.OperatorOuterClass.Operator +import org.apache.comet.serde.QueryPlanSerde.supportedDataType import org.apache.comet.serde.operator.CometSink case class CometLocalTableScanExec( @@ -131,15 +132,20 @@ object CometLocalTableScanExec extends CometSink[LocalTableScanExec] with DataTy // downstream expression serdes (issue #4789). override protected def scanFieldType(dt: DataType): DataType = dt.asNullable - // ArrowWriter (used by RowArrowReader) handles NullType via Utils.toArrowType + NullWriter; - // other types off DataTypeSupport's allow list (TimeType, intervals, ...) have no ArrowWriter - // coverage and must fall back to Spark. + // RowArrowReader handles NullType and intervals, but not TimeType. Non-default string collations + // remain unsupported here, matching DataTypeSupport's existing local-scan boundary. override def isTypeSupported( dt: DataType, name: String, - fallbackReasons: ListBuffer[String]): Boolean = dt match { - case _: NullType | _: YearMonthIntervalType | _: DayTimeIntervalType => true - case _ => super.isTypeSupported(dt, name, fallbackReasons) + fallbackReasons: ListBuffer[String]): Boolean = { + val supported = supportedDataType( + dt, + allowComplex = true, + allowIntervals = true, + allowTimeType = false, + allowAnyStringType = false) + if (!supported) super.isTypeSupported(dt, name, fallbackReasons) + supported } override def convert( diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index 71d17846305..141e3e19522 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -41,7 +41,7 @@ import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeExec, ShuffleExchangeLike, ShuffleOrigin} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics, SQLShuffleReadMetricsReporter, SQLShuffleWriteMetricsReporter} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, DataType, DateType, DayTimeIntervalType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, NullType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType, YearMonthIntervalType} +import org.apache.spark.sql.types.{BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.spark.util.MutablePair import org.apache.spark.util.collection.unsafe.sort.{PrefixComparators, RecordComparator} @@ -403,30 +403,6 @@ object CometShuffleExchangeExec false } - /** - * Determine which data types are supported as data columns in native shuffle. - * - * Native shuffle relies on the Arrow IPC writer to serialize batches to disk, so it should - * support all types that Comet supports. - */ - def supportedSerializableDataType(dt: DataType): Boolean = dt match { - case _: BooleanType | _: ByteType | _: ShortType | _: IntegerType | _: LongType | - _: FloatType | _: DoubleType | _: StringType | _: BinaryType | _: TimestampType | - _: TimestampNTZType | _: DecimalType | _: DateType | _: NullType | - _: YearMonthIntervalType | _: DayTimeIntervalType => - true - case dt if isTimeType(dt) => - true - case StructType(fields) => - fields.nonEmpty && fields.forall(f => supportedSerializableDataType(f.dataType)) - case ArrayType(elementType, _) => - supportedSerializableDataType(elementType) - case MapType(keyType, valueType, _) => - supportedSerializableDataType(keyType) && supportedSerializableDataType(valueType) - case _ => - false - } - val reasons = scala.collection.mutable.ListBuffer.empty[String] if (!isCometNativeShuffleMode(s.conf)) { @@ -437,7 +413,10 @@ object CometShuffleExchangeExec val inputs = s.child.output for (input <- inputs) { - if (!supportedSerializableDataType(input.dataType)) { + if (!QueryPlanSerde.supportedDataType( + input.dataType, + allowComplex = true, + allowIntervals = true)) { reasons += s"unsupported shuffle data type ${input.dataType} for input $input" return reasons.toSeq } @@ -529,32 +508,6 @@ object CometShuffleExchangeExec */ private def columnarShuffleFailureReasons(s: ShuffleExchangeExec): Seq[String] = { - /** - * Determine which data types are supported as data columns in columnar shuffle. - * - * Comet columnar shuffle used native code to convert Spark unsafe rows to Arrow batches, see - * shuffle/row.rs - */ - def supportedSerializableDataType(dt: DataType): Boolean = dt match { - case _: BooleanType | _: ByteType | _: ShortType | _: IntegerType | _: LongType | - _: FloatType | _: DoubleType | _: StringType | _: BinaryType | _: TimestampType | - _: TimestampNTZType | _: DecimalType | _: DateType | _: NullType => - true - case dt if isTimeType(dt) => - true - case StructType(fields) => - fields.nonEmpty && fields.forall(f => supportedSerializableDataType(f.dataType)) && - // Java Arrow stream reader cannot work on duplicate field name - fields.map(f => f.name).distinct.length == fields.length && - fields.nonEmpty - case ArrayType(elementType, _) => - supportedSerializableDataType(elementType) - case MapType(keyType, valueType, _) => - supportedSerializableDataType(keyType) && supportedSerializableDataType(valueType) - case _ => - false - } - val reasons = scala.collection.mutable.ListBuffer.empty[String] if (!isCometJVMShuffleMode(s.conf)) { @@ -575,7 +528,11 @@ object CometShuffleExchangeExec val inputs = s.child.output for (input <- inputs) { - if (!supportedSerializableDataType(input.dataType)) { + if (!QueryPlanSerde.supportedDataType( + input.dataType, + allowComplex = true, + // Java Arrow stream reader cannot work on duplicate field names. + allowDuplicateStructFieldNames = false)) { reasons += s"unsupported shuffle data type ${input.dataType} for input $input" return reasons.toSeq } diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index f3c442ad295..221d1ff08a2 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -37,11 +37,39 @@ import org.apache.spark.sql.types._ import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus, isSpark41Plus, isSpark42Plus} import org.apache.comet.serde.{CometAttributeReference, CometKnownFloatingPointNormalized, Unsupported} +import org.apache.comet.serde.QueryPlanSerde.supportedDataType import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { import testImplicits._ + test("supportedDataType applies capabilities recursively") { + Seq[DataType](YearMonthIntervalType(), DayTimeIntervalType()).foreach { interval => + val nested = StructType( + Seq(StructField( + "i", + ArrayType(MapType(StringType, interval, valueContainsNull = true), containsNull = true), + nullable = true))) + + assert(!supportedDataType(nested, allowComplex = true)) + assert(!supportedDataType(nested, allowIntervals = true)) + assert(supportedDataType(nested, allowComplex = true, allowIntervals = true)) + } + + val duplicateFields = ArrayType( + StructType(Seq(StructField("i", IntegerType), StructField("i", IntegerType))), + containsNull = true) + assert(supportedDataType(duplicateFields, allowComplex = true)) + assert( + !supportedDataType( + duplicateFields, + allowComplex = true, + allowDuplicateStructFieldNames = false)) + + assert(supportedDataType(StringType, allowAnyStringType = false)) + assert(!supportedDataType(CharType(1), allowAnyStringType = false)) + } + val ARITHMETIC_OVERFLOW_EXCEPTION_MSG = """[ARITHMETIC_OVERFLOW] integer overflow. If necessary set "spark.sql.ansi.enabled" to "false" to bypass this error""" val DIVIDE_BY_ZERO_EXCEPTION_MSG = From db5c3362a4826908e364cc0faded7cb0f1d6130b Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 28 Aug 2026 01:57:44 +0800 Subject: [PATCH 4/6] fix: keep columnar shuffle rejecting CalendarIntervalType The refactor placed CalendarIntervalType in the unconditional base arm of supportedDataType, which made columnar shuffle newly accept it. The native unsafe-row-to-Arrow converter (spark_unsafe/row.rs) has no interval support, so this would panic at runtime instead of falling back to Spark. Add an allowCalendarInterval capability (default true, matching every other pre-refactor call site) and pass false at the columnar shuffle boundary. Also document that the local-scan super call only records fallback reasons. Co-Authored-By: Claude Fable 5 --- .../scala/org/apache/comet/serde/QueryPlanSerde.scala | 5 ++++- .../spark/sql/comet/CometLocalTableScanExec.scala | 2 ++ .../execution/shuffle/CometShuffleExchangeExec.scala | 3 +++ .../scala/org/apache/comet/CometExpressionSuite.scala | 11 +++++++++++ 4 files changed, 20 insertions(+), 1 deletion(-) 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 7547a0818bf..4d0adad58b2 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -528,13 +528,16 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { dt: DataType, allowComplex: Boolean = false, allowIntervals: Boolean = false, + allowCalendarInterval: Boolean = true, allowTimeType: Boolean = true, allowAnyStringType: Boolean = true, allowDuplicateStructFieldNames: Boolean = true): Boolean = { def supported(dt: DataType): Boolean = dt match { case _: ByteType | _: ShortType | _: IntegerType | _: LongType | _: FloatType | _: DoubleType | _: BinaryType | _: TimestampType | _: TimestampNTZType | - _: DecimalType | _: DateType | _: BooleanType | _: NullType | CalendarIntervalType => + _: DecimalType | _: DateType | _: BooleanType | _: NullType => + true + case CalendarIntervalType if allowCalendarInterval => true case st: StringType if allowAnyStringType || st == StringType => true diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala index f28647997f8..acc43ca8eb1 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTableScanExec.scala @@ -144,6 +144,8 @@ object CometLocalTableScanExec extends CometSink[LocalTableScanExec] with DataTy allowIntervals = true, allowTimeType = false, allowAnyStringType = false) + // DataTypeSupport accepts no type that supportedDataType rejects here, so the super call + // cannot widen the accepted set; it only records the fallback reason for rejected types. if (supported) true else super.isTypeSupported(dt, name, fallbackReasons) } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index 141e3e19522..77b831e2920 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -531,6 +531,9 @@ object CometShuffleExchangeExec if (!QueryPlanSerde.supportedDataType( input.dataType, allowComplex = true, + // The native row-to-Arrow converter (spark_unsafe/row.rs) has no CalendarInterval + // support, so calendar intervals must fall back to Spark shuffle. + allowCalendarInterval = false, // Java Arrow stream reader cannot work on duplicate field names. allowDuplicateStructFieldNames = false)) { reasons += s"unsupported shuffle data type ${input.dataType} for input $input" diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 0175d8c9130..cf07abf7f4d 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -65,6 +65,17 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { allowComplex = true, allowDuplicateStructFieldNames = false)) + assert(supportedDataType(CalendarIntervalType)) + assert(!supportedDataType(CalendarIntervalType, allowCalendarInterval = false)) + val nestedCalendarInterval = + StructType(Seq(StructField("i", ArrayType(CalendarIntervalType, containsNull = true)))) + assert(supportedDataType(nestedCalendarInterval, allowComplex = true)) + assert( + !supportedDataType( + nestedCalendarInterval, + allowComplex = true, + allowCalendarInterval = false)) + assert(supportedDataType(StringType, allowAnyStringType = false)) assert(!supportedDataType(CharType(1), allowAnyStringType = false)) From d7f37b13842e299ef27b238b848518cfa9216cca Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 2 Sep 2026 12:15:29 +0800 Subject: [PATCH 5/6] test: cover data type boundary policies --- .../apache/comet/serde/QueryPlanSerde.scala | 29 +++++ .../apache/comet/CometExpressionSuite.scala | 45 -------- .../exec/CometColumnarShuffleSuite.scala | 25 ++++- .../comet/serde/QueryPlanSerdeSuite.scala | 104 ++++++++++++++++++ 4 files changed, 157 insertions(+), 46 deletions(-) create mode 100644 spark/src/test/scala/org/apache/comet/serde/QueryPlanSerdeSuite.scala 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 4ec84fc1e94..be089d6be93 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -547,6 +547,35 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { builder.build() } + /** + * Returns whether `dt` is supported at a caller's data-type boundary. + * + * The defaults preserve expression-serde behavior: primitive types, `CalendarIntervalType`, + * `TimeType`, and all `StringType` variants are accepted, while complex and ANSI interval types + * are rejected. Sinks and native shuffle enable complex and ANSI interval types because their + * Arrow IPC paths support them. Local scans additionally reject `TimeType` and non-default + * strings, while JVM columnar shuffle rejects ANSI intervals, calendar intervals, and duplicate + * struct field names because its unsafe-row-to-Arrow path cannot handle them. + * + * Note that the option polarity is mixed: `allowComplex` and `allowIntervals` are restrictive + * by default; the other four options are permissive by default. + * + * @param dt + * data type to check + * @param allowComplex + * recursively allow non-empty structs, arrays, and maps + * @param allowIntervals + * allow year-month and day-time interval types + * @param allowCalendarInterval + * allow calendar interval types + * @param allowTimeType + * allow Spark `TimeType` + * @param allowAnyStringType + * allow non-default `StringType` variants such as collated strings; when false, only the + * default `StringType` is accepted + * @param allowDuplicateStructFieldNames + * allow duplicate field names in nested structs + */ def supportedDataType( dt: DataType, allowComplex: Boolean = false, diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index f6102766537..25b43d9fe6e 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -34,56 +34,11 @@ import org.apache.spark.sql.internal.SQLConf.SESSION_LOCAL_TIMEZONE import org.apache.spark.sql.types._ import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus, isSpark41Plus, isSpark42Plus} -import org.apache.comet.serde.QueryPlanSerde.supportedDataType import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { import testImplicits._ - test("supportedDataType applies capabilities recursively") { - Seq[DataType](YearMonthIntervalType(), DayTimeIntervalType()).foreach { interval => - val nested = StructType( - Seq(StructField( - "i", - ArrayType(MapType(StringType, interval, valueContainsNull = true), containsNull = true), - nullable = true))) - - assert(!supportedDataType(nested, allowComplex = true)) - assert(!supportedDataType(nested, allowIntervals = true)) - assert(supportedDataType(nested, allowComplex = true, allowIntervals = true)) - } - - val duplicateFields = ArrayType( - StructType(Seq(StructField("i", IntegerType), StructField("i", IntegerType))), - containsNull = true) - assert(supportedDataType(duplicateFields, allowComplex = true)) - assert( - !supportedDataType( - duplicateFields, - allowComplex = true, - allowDuplicateStructFieldNames = false)) - - assert(supportedDataType(CalendarIntervalType)) - assert(!supportedDataType(CalendarIntervalType, allowCalendarInterval = false)) - val nestedCalendarInterval = - StructType(Seq(StructField("i", ArrayType(CalendarIntervalType, containsNull = true)))) - assert(supportedDataType(nestedCalendarInterval, allowComplex = true)) - assert( - !supportedDataType( - nestedCalendarInterval, - allowComplex = true, - allowCalendarInterval = false)) - - assert(supportedDataType(StringType, allowAnyStringType = false)) - assert(!supportedDataType(CharType(1), allowAnyStringType = false)) - - if (isSpark41Plus) { - val timeType = DataType.fromDDL("TIME") - assert(supportedDataType(timeType)) - assert(!supportedDataType(timeType, allowTimeType = false)) - } - } - val ARITHMETIC_OVERFLOW_EXCEPTION_MSG = """[ARITHMETIC_OVERFLOW] integer overflow. If necessary set "spark.sql.ansi.enabled" to "false" to bypass this error""" val DIVIDE_BY_ZERO_EXCEPTION_MSG = diff --git a/spark/src/test/scala/org/apache/comet/exec/CometColumnarShuffleSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometColumnarShuffleSuite.scala index bf354356386..911062dd33e 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometColumnarShuffleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometColumnarShuffleSuite.scala @@ -31,7 +31,7 @@ import org.apache.spark.{Partitioner, SparkConf} import org.apache.spark.sql.{CometTestBase, DataFrame, Row} import org.apache.spark.sql.comet.execution.shuffle.{CometShuffleDependency, CometShuffleExchangeExec, CometShuffleManager} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEShuffleReadExec, ShuffleQueryStageExec} -import org.apache.spark.sql.execution.exchange.ReusedExchangeExec +import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} import org.apache.spark.sql.execution.joins.SortMergeJoinExec import org.apache.spark.sql.functions.col import org.apache.spark.sql.internal.SQLConf @@ -75,6 +75,29 @@ abstract class CometColumnarShuffleSuite extends CometTestBase with AdaptiveSpar """.stripMargin).select($"r.*") checkSparkAnswer(df) + checkCometExchange(df, 0, false) + } + + test("Fallback to Spark when shuffling CalendarIntervalType data") { + val df = spark + .sql("select id, make_interval(1,2,3,4,5,6,7) as i from range(100)") + .repartition(4) + + assert(df.collect().length == 100) + + val plan = df.queryExecution.executedPlan + assert( + find(plan) { + case _: ShuffleExchangeExec => true + case _ => false + }.nonEmpty, + plan) + assert( + find(plan) { + case _: CometShuffleExchangeExec => true + case _ => false + }.isEmpty, + plan) } test("Unsupported types for SinglePartition should fallback to Spark") { diff --git a/spark/src/test/scala/org/apache/comet/serde/QueryPlanSerdeSuite.scala b/spark/src/test/scala/org/apache/comet/serde/QueryPlanSerdeSuite.scala new file mode 100644 index 00000000000..e973bbdee34 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/serde/QueryPlanSerdeSuite.scala @@ -0,0 +1,104 @@ +/* + * 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. + */ + +package org.apache.comet.serde + +import org.scalatest.funsuite.AnyFunSuite + +import org.apache.spark.sql.types._ + +import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus, isSpark41Plus} +import org.apache.comet.serde.QueryPlanSerde.supportedDataType + +class QueryPlanSerdeSuite extends AnyFunSuite { + + test("supportedDataType matches each caller boundary") { + val complex = ArrayType(IntegerType) + val nestedInterval = StructType( + Seq( + StructField( + "i", + ArrayType( + MapType(StringType, DayTimeIntervalType(), valueContainsNull = true), + containsNull = true)))) + val nestedCalendarInterval = + StructType(Seq(StructField("i", ArrayType(CalendarIntervalType, containsNull = true)))) + val duplicateFields = + ArrayType(StructType(Seq(StructField("i", IntegerType), StructField("i", IntegerType)))) + val emptyStruct = StructType(Nil) + val timeTypes = if (isSpark41Plus) Seq(DataType.fromDDL("TIME")) else Seq.empty + val collatedStrings = + if (isSpark40Plus) Seq(DataType.fromDDL("STRING COLLATE UTF8_LCASE")) else Seq.empty + + val boundaries: Seq[(String, DataType => Boolean, Seq[DataType], Seq[DataType])] = Seq( + ( + "expression serde defaults", + supportedDataType(_), + Seq(IntegerType, StringType, CalendarIntervalType) ++ timeTypes ++ collatedStrings, + Seq(complex, YearMonthIntervalType(), DayTimeIntervalType(), emptyStruct)), + ( + "CometSink", + supportedDataType(_, allowComplex = true, allowIntervals = true), + Seq(complex, nestedInterval, nestedCalendarInterval, duplicateFields) ++ + timeTypes ++ collatedStrings, + Seq(emptyStruct)), + ( + "CometLocalTableScanExec", + supportedDataType( + _, + allowComplex = true, + allowIntervals = true, + allowTimeType = false, + allowAnyStringType = false), + Seq( + IntegerType, + StringType, + complex, + nestedInterval, + nestedCalendarInterval, + duplicateFields), + Seq(emptyStruct) ++ timeTypes ++ collatedStrings), + ( + "native shuffle", + supportedDataType(_, allowComplex = true, allowIntervals = true), + Seq(complex, nestedInterval, nestedCalendarInterval, duplicateFields) ++ + timeTypes ++ collatedStrings, + Seq(emptyStruct)), + ( + "JVM columnar shuffle", + supportedDataType( + _, + allowComplex = true, + allowCalendarInterval = false, + allowDuplicateStructFieldNames = false), + Seq(IntegerType, StringType, complex) ++ timeTypes ++ collatedStrings, + Seq( + YearMonthIntervalType(), + DayTimeIntervalType(), + CalendarIntervalType, + nestedCalendarInterval, + duplicateFields, + emptyStruct))) + + boundaries.foreach { case (name, supports, accepted, rejected) => + accepted.foreach(dt => assert(supports(dt), s"$name should accept $dt")) + rejected.foreach(dt => assert(!supports(dt), s"$name should reject $dt")) + } + } +} From 86839b298d597a3f03142b4c8f1a5efde12a7b67 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 2 Sep 2026 12:29:26 +0800 Subject: [PATCH 6/6] ci: register QueryPlanSerdeSuite --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + 2 files changed, 2 insertions(+) diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 11a5619a4d2..16520c98ae6 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -466,6 +466,7 @@ jobs: org.apache.comet.CometWidthBucketSuite org.apache.comet.CometUuidExpressionSuite org.apache.comet.serde.CometScalarFunctionSuite + org.apache.comet.serde.QueryPlanSerdeSuite org.apache.comet.CometFallbackInvarianceSuite fail-fast: false name: ${{ matrix.profile.name }} [${{ matrix.suite.name }}] diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 7b7fae8fbea..19fb1423e02 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -239,6 +239,7 @@ jobs: org.apache.comet.CometWidthBucketSuite org.apache.comet.CometUuidExpressionSuite org.apache.comet.serde.CometScalarFunctionSuite + org.apache.comet.serde.QueryPlanSerdeSuite org.apache.comet.CometFallbackInvarianceSuite fail-fast: false