diff --git a/native/spark-expr/src/conversion_funcs/cast.rs b/native/spark-expr/src/conversion_funcs/cast.rs index b1b71e8267d..9e504486ca6 100644 --- a/native/spark-expr/src/conversion_funcs/cast.rs +++ b/native/spark-expr/src/conversion_funcs/cast.rs @@ -335,6 +335,7 @@ pub(crate) fn cast_array( } (Utf8View, Utf8) => Ok(cast_with_options(&array, to_type, &CAST_OPTIONS)?), (Struct(_), Utf8) => Ok(casts_struct_to_string(array.as_struct(), cast_options)?), + (Map(_, _), Utf8) => Ok(cast_map_to_string(array.as_map(), cast_options)?), (Struct(_), Struct(_)) => Ok(cast_struct_to_struct( array.as_struct(), &from_type, @@ -667,6 +668,68 @@ fn casts_struct_to_string( Ok(Arc::new(builder.finish())) } +fn cast_map_to_string( + array: &MapArray, + spark_cast_options: &SparkCastOptions, +) -> DataFusionResult { + let mut builder = StringBuilder::with_capacity(array.len(), array.len() * 16); + let mut str = String::with_capacity(array.len() * 16); + + let casted_keys = cast_array( + Arc::clone(array.keys()), + &DataType::Utf8, + spark_cast_options, + )?; + let casted_values = cast_array( + Arc::clone(array.values()), + &DataType::Utf8, + spark_cast_options, + )?; + let key_values = casted_keys + .as_any() + .downcast_ref::() + .expect("Casted keys should be StringArray"); + let value_values = casted_values + .as_any() + .downcast_ref::() + .expect("Casted values should be StringArray"); + + let offsets = array.offsets(); + for row_index in 0..array.len() { + if array.is_null(row_index) { + builder.append_null(); + } else { + str.clear(); + let start = offsets[row_index] as usize; + let end = offsets[row_index + 1] as usize; + + str.push('{'); + let mut first = true; + for idx in start..end { + if !first { + str.push_str(", "); + } + if key_values.is_null(idx) { + str.push_str(&spark_cast_options.null_string); + } else { + str.push_str(key_values.value(idx)); + } + str.push_str(" -> "); + if value_values.is_null(idx) { + str.push_str(&spark_cast_options.null_string); + } else { + str.push_str(value_values.value(idx)); + } + first = false; + } + str.push('}'); + builder.append_value(&str); + } + } + + Ok(Arc::new(builder.finish())) +} + impl Display for Cast { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { write!( @@ -847,7 +910,12 @@ fn cast_binary_to_string( #[cfg(test)] mod tests { use super::*; - use arrow::array::{BinaryArray, ListArray, NullArray, PrimitiveArray, StringArray}; + use arrow::array::builder::{ + Int32Builder, MapBuilder, StringBuilder, TimestampMicrosecondBuilder, + }; + use arrow::array::{ + BinaryArray, ListArray, MapFieldNames, NullArray, PrimitiveArray, StringArray, + }; use arrow::buffer::OffsetBuffer; use arrow::datatypes::{Field, Fields, Int32Type, TimestampMicrosecondType}; @@ -1156,6 +1224,92 @@ mod tests { } } + #[test] + fn test_cast_map_to_utf8() { + let mut map_builder = MapBuilder::new( + Some(MapFieldNames { + entry: "entries".into(), + key: "key".into(), + value: "value".into(), + }), + StringBuilder::new(), + Int32Builder::new(), + ); + + map_builder.keys().append_value("a"); + map_builder.values().append_value(1); + map_builder.keys().append_value("b"); + map_builder.values().append_null(); + map_builder.append(true).unwrap(); + + map_builder.append(true).unwrap(); + map_builder.append(false).unwrap(); + + let map_array: ArrayRef = Arc::new(map_builder.finish()); + let string_array = cast_array( + map_array, + &DataType::Utf8, + &SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + ) + .unwrap(); + let string_array = string_array.as_string::(); + assert_eq!(3, string_array.len()); + assert_eq!(r#"{a -> 1, b -> null}"#, string_array.value(0)); + assert_eq!(r#"{}"#, string_array.value(1)); + assert!(string_array.is_null(2)); + } + + #[test] + fn test_cast_map_to_utf8_ignores_values_outside_slice() { + let mut map_builder = MapBuilder::new( + None, + StringBuilder::new(), + TimestampMicrosecondBuilder::new(), + ); + map_builder.keys().append_value("hidden"); + map_builder.values().append_value(i64::MAX); + map_builder.append(true).unwrap(); + map_builder.keys().append_value("visible"); + map_builder.values().append_value(0); + map_builder.append(true).unwrap(); + + let string_array = cast_array( + Arc::new(map_builder.finish().slice(1, 1)), + &DataType::Utf8, + &SparkCastOptions::new(EvalMode::Ansi, "UTC", false), + ) + .unwrap(); + assert_eq!( + "{visible -> 1970-01-01 00:00:00}", + string_array.as_string::().value(0) + ); + } + + #[test] + fn test_cast_map_to_utf8_ignores_values_under_null_row() { + let mut map_builder = MapBuilder::new( + None, + StringBuilder::new(), + TimestampMicrosecondBuilder::new(), + ); + map_builder.keys().append_value("hidden"); + map_builder.values().append_value(i64::MAX); + map_builder.append(false).unwrap(); + map_builder.keys().append_value("visible"); + map_builder.values().append_value(0); + map_builder.append(true).unwrap(); + + let string_array = cast_array( + Arc::new(map_builder.finish()), + &DataType::Utf8, + &SparkCastOptions::new(EvalMode::Ansi, "UTC", false), + ) + .unwrap(); + let string_array = string_array.as_string::(); + assert!(string_array.is_null(0)); + assert_eq!("{visible -> 1970-01-01 00:00:00}", string_array.value(1)); + } + #[test] fn test_cast_string_array_to_string() { let values_array = diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala index c29300f9ed6..1d682cd0eb0 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala @@ -83,6 +83,7 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come * supported when their children are. */ def isSupportedDataType(dt: DataType): Boolean = dt match { + case NullType => true case BooleanType | ByteType | ShortType | IntegerType | LongType => true case FloatType | DoubleType => true case _: DecimalType => true 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..d820993bf91 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenInput.scala @@ -603,6 +603,7 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { * element or struct field. `idx` is the index/ordinal token (e.g. `"__i"` or `"3"`). */ private def elementGetterCall(dt: DataType, idx: String): String = dt match { + case NullType => "null" case BooleanType => s"getBoolean($idx)" case ByteType => s"getByte($idx)" case ShortType => s"getShort($idx)" @@ -691,6 +692,8 @@ private[codegen] object CometBatchKernelCodegenInput extends CometTypeShim { if (elementNullable) " if (isNullAt(i)) return null;\n" else "" elemType match { + case NullType => + "" case BooleanType => s""" @Override | public boolean getBoolean(int i) { 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 33e6c0c0355..71998c9e836 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -159,6 +159,7 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { /** Concrete Arrow vector class name for the output type, used to cast `outRaw` once. */ private def outputVectorClass(dataType: DataType): String = dataType match { + case NullType => classOf[NullVector].getName case BooleanType => classOf[BitVector].getName case ByteType => classOf[TinyIntVector].getName case ShortType => classOf[SmallIntVector].getName @@ -209,6 +210,8 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { dataType: DataType, ctx: CodegenContext, nested: Boolean = false): OutputEmit = dataType match { + case NullType => + OutputEmit("", "") case BooleanType => val set = if (nested) "setSafe" else "set" OutputEmit("", s"$targetVec.$set($idx, $source ? 1 : 0);") @@ -407,6 +410,7 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { */ private def emitSpecializedGetterExpr(target: String, idx: String, elemType: DataType): String = elemType match { + case NullType => "null" case BooleanType => s"$target.getBoolean($idx)" case ByteType => s"$target.getByte($idx)" case ShortType => s"$target.getShort($idx)" diff --git a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala index d29ef7cd3b7..e4f33986d56 100644 --- a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala +++ b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala @@ -321,6 +321,8 @@ object CometCast Compatible() case DataTypes.BinaryType => Compatible() + case DataTypes.NullType => + Compatible() case StructType(fields) => for (field <- fields) { isSupported(field.dataType, DataTypes.StringType, timeZoneId, evalMode) match { @@ -332,6 +334,13 @@ object CometCast } } Compatible() + case MapType(keyType, valueType, _) => + isSupported(keyType, DataTypes.StringType, timeZoneId, evalMode) match { + case Compatible(_, _) => + isSupported(valueType, DataTypes.StringType, timeZoneId, evalMode) + case other => + other + } case _ => unsupported(fromType, DataTypes.StringType) } } diff --git a/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string.sql b/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string.sql index dd42db8e083..b7cc364799c 100644 --- a/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string.sql +++ b/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string.sql @@ -149,6 +149,7 @@ SELECT cast(named_struct('a', named_struct('b', named_struct('c', 1, 'd', 'leaf' query SELECT cast(named_struct('s1', '', 's2', ' ', 's3', cast(null as string)) as string) +-- Map-valued field: supported via recursive map -> string casting. query SELECT cast(named_struct('m', map('k', 1)) as string) @@ -269,15 +270,14 @@ SELECT cast(array(cast(1.5 as double), cast('NaN' as double), cast('-Infinity' a query SELECT cast(array(array(array(1, 2), array(3)), array(array(cast(null as int)))) as string) --- Array of map: map-to-string is routed through the codegen dispatcher via the outer array. +-- Array of map: supported via recursive map -> string casting. query SELECT cast(array(map('k', 1)) as string) -- ---------------------------------------------------------------------------- -- Map → string -- ---------------------------------------------------------------------------- --- Comet has no native map-to-string cast; `CometCast` mixes in `CodegenDispatchFallback`, so --- these stay native via the codegen dispatcher and match Spark exactly. +-- Comet now implements map-to-string casts, including nested maps. -- Note: maps materialized through parquet have nondeterministic entry order, so map column -- tests use literal maps directly rather than reading from a parquet table. diff --git a/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string_legacy.sql b/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string_legacy.sql index fcdb94b3526..bd1a5d82231 100644 --- a/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string_legacy.sql +++ b/spark/src/test/resources/sql-tests/expressions/cast/cast_complex_types_to_string_legacy.sql @@ -34,6 +34,10 @@ SELECT CAST(array(1, 2, null) AS STRING) query SELECT CAST(map('a', 1, 'b', null) AS STRING) +-- Empty Map → string. +query +SELECT CAST(map() AS STRING) + -- Nested complex types via the outer struct. query SELECT CAST(struct(array(1, null), map('k', null)) AS STRING) diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 5806cb35015..96575d2a346 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -22,10 +22,11 @@ package org.apache.comet import scala.util.Random import org.apache.arrow.vector._ +import org.apache.arrow.vector.complex.ListVector 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.catalyst.expressions.{BoundReference, Cast, Coalesce, CreateArray, CreateMap, CreateNamedStruct, Expression, GetArrayItem, IsNull, Literal, MapConcat} import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ @@ -33,7 +34,7 @@ import org.apache.spark.unsafe.types.UTF8String import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus import org.apache.comet.codegen.CometBatchKernelCodegen -import org.apache.comet.codegen.CometBatchKernelCodegen.ArrowColumnSpec +import org.apache.comet.codegen.CometBatchKernelCodegen.{ArrayColumnSpec, ArrowColumnSpec, ScalarColumnSpec} import org.apache.comet.udf.codegen.CometScalaUDFCodegen import org.apache.comet.vector.CometVector @@ -81,6 +82,96 @@ class CometCodegenSuite } } + test("codegen dispatch safely casts a NullType column to a primitive") { + val input = new NullVector("in", 2) + val expr = Coalesce( + Seq( + Cast(BoundReference(0, NullType, nullable = true), IntegerType, ansiEnabled = true), + Literal(42))) + val field = CometBatchKernelCodegen.toFfiArrowField("out", IntegerType, nullable = false) + val output = CometBatchKernelCodegen.allocateOutput(field, 2, 0) + try { + input.setValueCount(2) + val spec = ArrowColumnSpec(classOf[NullVector], nullable = true) + val kernel = CometBatchKernelCodegen.compile(expr, IndexedSeq(spec)).newInstance() + kernel.init(0) + kernel.process(Array(input), output, 2) + output.setValueCount(2) + + val comet = CometVector.getVector(output, null) + assert(comet.getInt(0) === 42) + assert(comet.getInt(1) === 42) + } finally { + output.close() + input.close() + } + } + + test("codegen dispatch reads an Array element in a primitive expression") { + val arrayType = ArrayType(NullType, containsNull = true) + val inputField = CometBatchKernelCodegen.toFfiArrowField("in", arrayType, nullable = false) + val input = CometBatchKernelCodegen + .allocateOutput(inputField, 2, 0) + .asInstanceOf[ListVector] + val expr = IsNull( + GetArrayItem( + BoundReference(0, arrayType, nullable = false), + Literal(0), + failOnError = false)) + val outputField = + CometBatchKernelCodegen.toFfiArrowField("out", BooleanType, nullable = false) + val output = CometBatchKernelCodegen.allocateOutput(outputField, 2, 0) + try { + input.startNewValue(0) + input.endValue(0, 1) + input.startNewValue(1) + input.endValue(1, 1) + input.getDataVector.asInstanceOf[NullVector].setValueCount(2) + input.setValueCount(2) + + val spec = ArrayColumnSpec( + nullable = false, + elementSparkType = NullType, + element = ScalarColumnSpec(classOf[NullVector], nullable = true)) + val kernel = CometBatchKernelCodegen.compile(expr, IndexedSeq(spec)).newInstance() + kernel.init(0) + kernel.process(Array(input), output, 2) + output.setValueCount(2) + + val comet = CometVector.getVector(output, null) + assert(comet.getBoolean(0)) + assert(comet.getBoolean(1)) + } finally { + output.close() + input.close() + } + } + + test("codegen dispatch writes Array output") { + val input = new NullVector("in", 2) + val expr = CreateArray(Seq(BoundReference(0, NullType, nullable = true))) + val outputField = + CometBatchKernelCodegen.toFfiArrowField("out", expr.dataType, nullable = false) + val output = CometBatchKernelCodegen.allocateOutput(outputField, 2, 0) + try { + input.setValueCount(2) + val spec = ArrowColumnSpec(classOf[NullVector], nullable = true) + val kernel = CometBatchKernelCodegen.compile(expr, IndexedSeq(spec)).newInstance() + kernel.init(0) + kernel.process(Array(input), output, 2) + output.setValueCount(2) + + val comet = CometVector.getVector(output, null) + val first = comet.getArray(0) + val second = comet.getArray(1) + assert(first.numElements() === 1 && first.isNullAt(0)) + assert(second.numElements() === 1 && second.isNullAt(0)) + } finally { + output.close() + input.close() + } + } + test("codegen kernel round-trips CalendarIntervalType") { val input = new IntervalMonthDayNanoVector("in", CometArrowAllocator) val field = diff --git a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala index 9bfa8f63774..e9df9bb7521 100644 --- a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala @@ -34,7 +34,7 @@ import org.apache.spark.sql.catalyst.parser.ParseException import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.functions.{col, monotonically_increasing_id} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, DataType, DataTypes, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, ShortType, StringType, StructField, StructType, TimestampType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, CalendarIntervalType, DataType, DataTypes, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, ShortType, StringType, StructField, StructType, TimestampType} import org.apache.comet.expressions.{CometCast, CometEvalMode} import org.apache.comet.rules.CometScanTypeChecker @@ -1877,18 +1877,36 @@ class CometNativeCastSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("cast MapType propagates Unsupported from nested value cast") { + test("cast MapType to StringType is Compatible") { + val fromType = MapType(IntegerType, IntegerType) + assert( + CometCast.isSupported(fromType, DataTypes.StringType, None, CometEvalMode.LEGACY) == + Compatible()) + } + + test("cast MapType propagates supported nested value cast") { // Map> → Map: the inner Map → String - // cast is Unsupported, and that must propagate through the outer Map - // arm rather than being silently swallowed. + // cast is now supported and must propagate through the outer Map arm. val innerFrom = MapType(IntegerType, IntegerType) - val expectedMessage = s"Cast from $innerFrom to ${DataTypes.StringType} is not supported" assert( CometCast.isSupported( MapType(IntegerType, innerFrom), MapType(IntegerType, StringType), None, CometEvalMode.LEGACY) == + Compatible()) + } + + test("cast MapType propagates Unsupported from nested value cast") { + val unsupportedValueType = CalendarIntervalType + val expectedMessage = + s"Cast from $unsupportedValueType to ${DataTypes.StringType} is not supported" + assert( + CometCast.isSupported( + MapType(IntegerType, unsupportedValueType), + DataTypes.StringType, + None, + CometEvalMode.LEGACY) == Unsupported(Some(expectedMessage))) }