diff --git a/sql/hive/src/main/scala/org/apache/spark/sql/hive/HiveInspectors.scala b/sql/hive/src/main/scala/org/apache/spark/sql/hive/HiveInspectors.scala index 76d9a08f603d0..b8758407877f0 100644 --- a/sql/hive/src/main/scala/org/apache/spark/sql/hive/HiveInspectors.scala +++ b/sql/hive/src/main/scala/org/apache/spark/sql/hive/HiveInspectors.scala @@ -27,7 +27,7 @@ import org.apache.hadoop.hive.common.`type`.{HiveChar, HiveDecimal, HiveInterval import org.apache.hadoop.hive.serde2.{io => hiveIo} import org.apache.hadoop.hive.serde2.objectinspector.{StructField => HiveStructField, _} import org.apache.hadoop.hive.serde2.objectinspector.primitive._ -import org.apache.hadoop.hive.serde2.typeinfo.{DecimalTypeInfo, TypeInfoFactory} +import org.apache.hadoop.hive.serde2.typeinfo.{CharTypeInfo, DecimalTypeInfo, TypeInfoFactory, VarcharTypeInfo} import org.apache.spark.SparkException import org.apache.spark.sql.AnalysisException @@ -36,6 +36,7 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util._ import org.apache.spark.sql.errors.DataTypeErrors.toSQLType import org.apache.spark.sql.execution.datasources.DaysWritable +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.TimestampNanosVal @@ -103,10 +104,6 @@ import org.apache.spark.unsafe.types.UTF8String * Struct: Object[] / java.util.List / java POJO * Union: class StandardUnion { byte tag; Object object } * - * NOTICE: HiveVarchar/HiveChar is not supported by catalyst, it will be simply considered as - * String type. - * - * * 2. Hive ObjectInspector is a group of flexible APIs to inspect value in different data * representation, and developers can extend those API as needed, so technically, * object inspector supports arbitrary data type in java. @@ -280,7 +277,30 @@ private[hive] trait HiveInspectors { (o: Any) => x.getWritableConstantValue case x: PrimitiveObjectInspector => x match { - // TODO we don't support the HiveVarcharObjectInspector yet. + case hvoi: HiveVarcharObjectInspector => + val length = hvoi.getTypeInfo.asInstanceOf[VarcharTypeInfo].getLength + charVarcharWrapper( + dataType match { + case v: VarcharType => Some(v.length) + case _ => None + }, + length, + hvoi.preferWritable(), + CharVarcharCodegenUtils.varcharTypeWriteSideCheck, + (value, size) => new HiveVarchar(value, size), + (value, size) => new hiveIo.HiveVarcharWritable(new HiveVarchar(value, size))) + case hcoi: HiveCharObjectInspector => + val length = hcoi.getTypeInfo.asInstanceOf[CharTypeInfo].getLength + charVarcharWrapper( + dataType match { + case c: CharType => Some(c.length) + case _ => None + }, + length, + hcoi.preferWritable(), + CharVarcharCodegenUtils.charTypeWriteSideCheck, + (value, size) => new HiveChar(value, size), + (value, size) => new hiveIo.HiveCharWritable(new HiveChar(value, size))) case _: StringObjectInspector if x.preferWritable() => withNullSafe(o => getStringWritable(o)) case _: StringObjectInspector => @@ -313,21 +333,6 @@ private[hive] trait HiveInspectors { withNullSafe(o => getByteWritable(o)) case _: ByteObjectInspector => withNullSafe(o => o.asInstanceOf[java.lang.Byte]) - // To spark HiveVarchar and HiveChar are same as string - case _: HiveVarcharObjectInspector if x.preferWritable() => - withNullSafe(o => getStringWritable(o)) - case _: HiveVarcharObjectInspector => - withNullSafe { o => - val s = o.asInstanceOf[UTF8String].toString - new HiveVarchar(s, s.length) - } - case _: HiveCharObjectInspector if x.preferWritable() => - withNullSafe(o => getStringWritable(o)) - case _: HiveCharObjectInspector => - withNullSafe { o => - val s = o.asInstanceOf[UTF8String].toString - new HiveChar(s, s.length) - } case _: JavaHiveDecimalObjectInspector => withNullSafe(o => HiveDecimal.create(o.asInstanceOf[Decimal].toJavaBigDecimal)) @@ -490,6 +495,39 @@ private[hive] trait HiveInspectors { identity[Any] } + private def charVarcharWrapper( + targetLength: Option[Int], + inspectorLength: Int, + preferWritable: Boolean, + writeSideCheck: (UTF8String, Int) => UTF8String, + toJava: (String, Int) => Any, + toWritable: (String, Int) => Any): Any => Any = { + if (preferWritable) { + targetLength match { + case Some(length) => + withNullSafe { value => + val checked = writeSideCheck(value.asInstanceOf[UTF8String], length) + toWritable(checked.toString, inspectorLength) + } + case None => + withNullSafe(value => getStringWritable(value)) + } + } else { + targetLength match { + case Some(length) => + withNullSafe { value => + val checked = writeSideCheck(value.asInstanceOf[UTF8String], length) + toJava(checked.toString, inspectorLength) + } + case None => + withNullSafe { value => + val string = value.asInstanceOf[UTF8String].toString + toJava(string, string.length) + } + } + } + } + /** * Builds unwrappers ahead of time according to object inspector * types to avoid pattern matching and branching costs per row. @@ -583,7 +621,8 @@ private[hive] trait HiveInspectors { val constant = ym.getWritableConstantValue.asInstanceOf[HiveIntervalYearMonth] _ => constant.getTotalMonths case pi: PrimitiveObjectInspector => pi match { - // We think HiveVarchar/HiveChar is also a String + // Untyped conversion extracts the raw value. The typed overload applies CHAR/VARCHAR + // padding and overflow checks. case hvoi: HiveVarcharObjectInspector if hvoi.preferWritable() => data: Any => { if (data != null) { @@ -806,8 +845,9 @@ private[hive] trait HiveInspectors { * Catalyst `dataType` to preserve nanosecond timestamp precision. The plain * `unwrapperFor(ObjectInspector)` cannot do this because a Hive `TimestampObjectInspector` * maps to micros by default; here the nanos timestamp types are produced as `TimestampNanosVal`, - * recursing through array/map/struct so nested nanos timestamps round-trip correctly. Any other - * type is delegated to the `ObjectInspector`-only overload. + * recursing through array/map/struct so nested nanos timestamps round-trip correctly. CHAR and + * VARCHAR targets also apply their read-side length and padding checks. Any other type is + * delegated to the `ObjectInspector`-only overload. */ def unwrapperFor(objectInspector: ObjectInspector, dataType: DataType): Any => Any = (objectInspector, dataType) match { @@ -829,6 +869,26 @@ private[hive] trait HiveInspectors { null } } + case (_, c: CharType) => + val unwrapper = unwrapperFor(objectInspector) + data: Any => { + val value = unwrapper(data).asInstanceOf[UTF8String] + if (value == null) { + null + } else { + CharVarcharCodegenUtils.charTypeReadSideCheck(value, c.length) + } + } + case (_, v: VarcharType) => + val unwrapper = unwrapperFor(objectInspector) + data: Any => { + val value = unwrapper(data).asInstanceOf[UTF8String] + if (value == null) { + null + } else { + CharVarcharCodegenUtils.varcharTypeReadSideCheck(value, v.length) + } + } case (li: ListObjectInspector, ArrayType(elementType, _)) => val unwrapper = unwrapperFor(li.getListElementObjectInspector, elementType) data: Any => { @@ -843,10 +903,24 @@ private[hive] trait HiveInspectors { case (mi: MapObjectInspector, MapType(keyType, valueType, _)) => val keyUnwrapper = unwrapperFor(mi.getMapKeyObjectInspector, keyType) val valueUnwrapper = unwrapperFor(mi.getMapValueObjectInspector, valueType) + val needsKeyValidation = CharVarcharUtils.hasCharVarchar(keyType) data: Any => { if (data != null) { val map = mi.getMap(data) - if (map == null) null else ArrayBasedMapData(map, keyUnwrapper, valueUnwrapper) + if (map == null) { + null + } else if (needsKeyValidation) { + val builder = new ArrayBasedMapBuilder(keyType, valueType) + val it = map.entrySet().iterator().asInstanceOf[ + java.util.Iterator[java.util.Map.Entry[Any, Any]]] + while (it.hasNext) { + val e = it.next() + builder.put(keyUnwrapper(e.getKey), valueUnwrapper(e.getValue)) + } + builder.build() + } else { + ArrayBasedMapData(map, keyUnwrapper, valueUnwrapper) + } } else { null } @@ -939,7 +1013,14 @@ private[hive] trait HiveInspectors { case MapType(keyType, valueType, _) => ObjectInspectorFactory.getStandardMapObjectInspector( toInspector(keyType), toInspector(valueType)) - case StringType => PrimitiveObjectInspectorFactory.javaStringObjectInspector + // Hive object inspectors preserve the CHAR/VARCHAR length but have no collation metadata. + case c: CharType => + PrimitiveObjectInspectorFactory.getPrimitiveJavaObjectInspector( + toHiveCharTypeInfo(c)) + case v: VarcharType => + PrimitiveObjectInspectorFactory.getPrimitiveJavaObjectInspector( + toHiveVarcharTypeInfo(v)) + case _: StringType => PrimitiveObjectInspectorFactory.javaStringObjectInspector case IntegerType => PrimitiveObjectInspectorFactory.javaIntObjectInspector case DoubleType => PrimitiveObjectInspectorFactory.javaDoubleObjectInspector case BooleanType => PrimitiveObjectInspectorFactory.javaBooleanObjectInspector @@ -978,6 +1059,20 @@ private[hive] trait HiveInspectors { messageParameters = Map("typeName" -> toSQLType(dataType))) } + private def toHiveCharTypeInfo(dataType: CharType): CharTypeInfo = { + if (dataType.length < 1 || dataType.length > HiveChar.MAX_CHAR_LENGTH) { + throw unsupportedHiveType(dataType) + } + TypeInfoFactory.getCharTypeInfo(dataType.length) + } + + private def toHiveVarcharTypeInfo(dataType: VarcharType): VarcharTypeInfo = { + if (dataType.length < 1 || dataType.length > HiveVarchar.MAX_VARCHAR_LENGTH) { + throw unsupportedHiveType(dataType) + } + TypeInfoFactory.getVarcharTypeInfo(dataType.length) + } + /** * Map the catalyst expression to ObjectInspector, however, * if the expression is `Literal` or foldable, a constant writable object inspector returns; @@ -986,7 +1081,11 @@ private[hive] trait HiveInspectors { * @return Hive java objectinspector (recursively). */ def toInspector(expr: Expression): ObjectInspector = expr match { - case Literal(value, StringType) => + case Literal(value, c: CharType) => + getHiveCharWritableConstantObjectInspector(value, c) + case Literal(value, v: VarcharType) => + getHiveVarcharWritableConstantObjectInspector(value, v) + case Literal(value, _: StringType) => getStringWritableConstantObjectInspector(value) case Literal(value, IntegerType) => getIntWritableConstantObjectInspector(value) @@ -1075,23 +1174,41 @@ private[hive] trait HiveInspectors { case _ => e.children.forall(canEarlyEval) } - def inspectorToDataType(inspector: ObjectInspector): DataType = inspector match { + def inspectorToDataType(inspector: ObjectInspector): DataType = + inspectorToDataType(inspector, SQLConf.get.charVarcharFirstClassTypes) + + /** + * Maps a Hive inspector to a Catalyst type. When `preserveCharVarchar` is true, CHAR/VARCHAR + * inspectors stay CHAR/VARCHAR even if first-class types are disabled. Runtime compatibility + * checks use that mode so analysis snapshots are not compared against a conf-dependent STRING + * rewrite. + */ + def inspectorToDataType( + inspector: ObjectInspector, + preserveCharVarchar: Boolean): DataType = inspector match { case s: StructObjectInspector => StructType(s.getAllStructFieldRefs.asScala.map(f => types.StructField( - f.getFieldName, inspectorToDataType(f.getFieldObjectInspector), nullable = true) + f.getFieldName, + inspectorToDataType(f.getFieldObjectInspector, preserveCharVarchar), + nullable = true) ).toArray) - case l: ListObjectInspector => ArrayType(inspectorToDataType(l.getListElementObjectInspector)) + case l: ListObjectInspector => + ArrayType(inspectorToDataType(l.getListElementObjectInspector, preserveCharVarchar)) case m: MapObjectInspector => MapType( - inspectorToDataType(m.getMapKeyObjectInspector), - inspectorToDataType(m.getMapValueObjectInspector)) + inspectorToDataType(m.getMapKeyObjectInspector, preserveCharVarchar), + inspectorToDataType(m.getMapValueObjectInspector, preserveCharVarchar)) case _: WritableStringObjectInspector => StringType case _: JavaStringObjectInspector => StringType - case _: WritableHiveVarcharObjectInspector => StringType - case _: JavaHiveVarcharObjectInspector => StringType - case _: WritableHiveCharObjectInspector => StringType - case _: JavaHiveCharObjectInspector => StringType + // Hive object inspectors cannot represent collations, so Hive function results use the + // default collation while preserving CHAR/VARCHAR length under first-class semantics. + case hvoi: HiveVarcharObjectInspector if preserveCharVarchar => + VarcharType(hvoi.getTypeInfo.asInstanceOf[VarcharTypeInfo].getLength) + case _: HiveVarcharObjectInspector => StringType + case hcoi: HiveCharObjectInspector if preserveCharVarchar => + CharType(hcoi.getTypeInfo.asInstanceOf[CharTypeInfo].getLength) + case _: HiveCharObjectInspector => StringType case _: WritableIntObjectInspector => IntegerType case _: JavaIntObjectInspector => IntegerType case _: WritableDoubleObjectInspector => DoubleType @@ -1122,6 +1239,51 @@ private[hive] trait HiveInspectors { case _: JavaVoidObjectInspector => NullType } + /** + * Analysis snapshots the Catalyst return type, but runtime inspectors are rebuilt from the + * current children (including foldability and session CHAR/VARCHAR settings). STRING may drift + * to or from a bounded string type across that boundary. Two bounded types must match exactly: + * accepting a different kind or length would apply the snapshotted conversion to an incompatible + * runtime value. + */ + def checkCompatibleHiveReturnType( + inspector: ObjectInspector, + expectedType: DataType): Unit = { + checkCompatibleHiveReturnType( + inspectorToDataType(inspector, preserveCharVarchar = true), + expectedType) + } + + def checkCompatibleHiveReturnType( + runtimeType: DataType, + expectedType: DataType): Unit = { + if (!compatibleHiveReturnType(runtimeType, expectedType)) { + throw SparkException.internalError( + s"Hive function runtime type ${runtimeType.catalogString} is incompatible " + + s"with analysis type ${expectedType.catalogString}.") + } + } + + private def compatibleHiveReturnType( + runtimeType: DataType, + expectedType: DataType): Boolean = { + (runtimeType, expectedType) match { + case (rt: CharType, et: CharType) => rt == et + case (rt: VarcharType, et: VarcharType) => rt == et + case (_: CharType | _: VarcharType, _: CharType | _: VarcharType) => false + case (_: StringType, _: CharType | _: VarcharType) => true + case (_: CharType | _: VarcharType, _: StringType) => true + case (ArrayType(rt, _), ArrayType(et, _)) => compatibleHiveReturnType(rt, et) + case (MapType(rk, rv, _), MapType(ek, ev, _)) => + compatibleHiveReturnType(rk, ek) && compatibleHiveReturnType(rv, ev) + case (rt: StructType, et: StructType) if rt.length == et.length => + rt.fields.zip(et.fields).forall { case (rf, ef) => + rf.name == ef.name && compatibleHiveReturnType(rf.dataType, ef.dataType) + } + case (rt, et) => rt.sameType(et) + } + } + private def decimalTypeInfoToCatalyst(inspector: PrimitiveObjectInspector): DecimalType = { val info = inspector.getTypeInfo.asInstanceOf[DecimalTypeInfo] DecimalType(info.precision(), info.scale()) @@ -1131,6 +1293,36 @@ private[hive] trait HiveInspectors { PrimitiveObjectInspectorFactory.getPrimitiveWritableConstantObjectInspector( TypeInfoFactory.stringTypeInfo, getStringWritable(value)) + private def getHiveCharWritableConstantObjectInspector( + value: Any, + dataType: CharType): ObjectInspector = { + val writable = if (value == null) { + null + } else { + val checked = CharVarcharCodegenUtils.charTypeWriteSideCheck( + value.asInstanceOf[UTF8String], dataType.length) + new hiveIo.HiveCharWritable( + new HiveChar(checked.toString, dataType.length)) + } + PrimitiveObjectInspectorFactory.getPrimitiveWritableConstantObjectInspector( + toHiveCharTypeInfo(dataType), writable) + } + + private def getHiveVarcharWritableConstantObjectInspector( + value: Any, + dataType: VarcharType): ObjectInspector = { + val writable = if (value == null) { + null + } else { + val checked = CharVarcharCodegenUtils.varcharTypeWriteSideCheck( + value.asInstanceOf[UTF8String], dataType.length) + new hiveIo.HiveVarcharWritable( + new HiveVarchar(checked.toString, dataType.length)) + } + PrimitiveObjectInspectorFactory.getPrimitiveWritableConstantObjectInspector( + toHiveVarcharTypeInfo(dataType), writable) + } + private def getIntWritableConstantObjectInspector(value: Any): ObjectInspector = PrimitiveObjectInspectorFactory.getPrimitiveWritableConstantObjectInspector( TypeInfoFactory.intTypeInfo, getIntWritable(value)) @@ -1296,7 +1488,9 @@ private[hive] trait HiveInspectors { case IntegerType => intTypeInfo case LongType => longTypeInfo case ShortType => shortTypeInfo - case StringType => stringTypeInfo + case c: CharType => toHiveCharTypeInfo(c) + case v: VarcharType => toHiveVarcharTypeInfo(v) + case _: StringType => stringTypeInfo case d: DecimalType => decimalTypeInfo(d) case DateType => dateTypeInfo case TimestampType => timestampTypeInfo diff --git a/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFEvaluators.scala b/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFEvaluators.scala index f09ee7eed93ef..a24767701376d 100644 --- a/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFEvaluators.scala +++ b/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFEvaluators.scala @@ -111,34 +111,36 @@ class HiveSimpleUDFEvaluator( } } -class HiveGenericUDFEvaluator( - funcWrapper: HiveFunctionWrapper, children: Seq[Expression]) - extends HiveUDFEvaluatorBase[GenericUDF](funcWrapper, children) { - - // SPARK-58792: copied expression nodes (e.g. via withNewChildrenInternal) share one - // HiveFunctionWrapper, whose cached GenericUDF instance is mutable: initialize() - // rewrites its converters and output holders based on the arguments of whichever - // copy initialized it last. Give every evaluator its own clone so copied nodes - // cannot corrupt each other. - @transient - override lazy val function: GenericUDF = - HiveFunctionRegistryUtils.cloneGenericUDF(funcWrapper.createFunction[GenericUDF]()) - - @transient - private lazy val argumentInspectors = children.map(toInspector).toArray +private[hive] object HiveGenericUDFEvaluator extends HiveInspectors { + + /** + * Driver-side Hive initialization for `SELECT hive_udf(...)`. Returns the Catalyst type + * (for example CHAR(5) from a CHAR inspector). `HiveGenericUDF.apply` stores it. + */ + def inferReturnType( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression]): DataType = { + val function = + HiveFunctionRegistryUtils.cloneGenericUDF(funcWrapper.createFunction[GenericUDF]()) + inspectorToDataType(initialize(function, children.map(toInspector).toArray)) + } - @transient - lazy val returnInspector = { + def initialize( + function: GenericUDF, + argumentInspectors: Array[ObjectInspector]): ObjectInspector = { // Inline o.a.h.hive.ql.udf.generic.GenericUDF#initializeAndFoldConstants, but // eliminate calls o.a.h.hive.ql.exec.FunctionRegistry to avoid initializing Hive // built-in UDFs. val oi = function.initialize(argumentInspectors) + val udfType = function.getClass.getAnnotation(classOf[HiveUDFType]) + val isDeterministic = + udfType != null && udfType.deterministic() && !udfType.stateful() // If the UDF depends on any external resources, we can't fold because the // resources may not be available at compile time. if (function.getRequiredFiles == null && function.getRequiredJars == null && argumentInspectors.forall(ObjectInspectorUtils.isConstantObjectInspector) && !ObjectInspectorUtils.isConstantObjectInspector(oi) && - isUDFDeterministic && + isDeterministic && ObjectInspectorUtils.supportsConstantObjectInspector(oi)) { val argumentValues: Array[DeferredObject] = argumentInspectors.map { argumentInspector => new GenericUDF.DeferredJavaObject( @@ -155,6 +157,32 @@ class HiveGenericUDFEvaluator( oi } } +} + +private[hive] class HiveGenericUDFEvaluator( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + catalystReturnType: DataType) + extends HiveUDFEvaluatorBase[GenericUDF](funcWrapper, children) { + + // SPARK-58792: copied expression nodes (e.g. via withNewChildrenInternal) share one + // HiveFunctionWrapper, whose cached GenericUDF instance is mutable: initialize() + // rewrites its converters and output holders based on the arguments of whichever + // copy initialized it last. Give every evaluator its own clone so copied nodes + // cannot corrupt each other. + @transient + override lazy val function: GenericUDF = + HiveFunctionRegistryUtils.cloneGenericUDF(funcWrapper.createFunction[GenericUDF]()) + + @transient + private lazy val argumentInspectors = children.map(toInspector).toArray + + @transient + lazy val returnInspector = { + val inspector = HiveGenericUDFEvaluator.initialize(function, argumentInspectors) + checkCompatibleHiveReturnType(inspector, catalystReturnType) + inspector + } @transient private lazy val deferredObjects: Array[DeferredObject] = argumentInspectors.zip(children).map { @@ -162,9 +190,9 @@ class HiveGenericUDFEvaluator( } @transient - private lazy val unwrapper: Any => Any = unwrapperFor(returnInspector) + private lazy val unwrapper: Any => Any = unwrapperFor(returnInspector, catalystReturnType) - override def returnType: DataType = inspectorToDataType(returnInspector) + override def returnType: DataType = catalystReturnType def setArg(index: Int, arg: Any): Unit = deferredObjects(index).asInstanceOf[DeferredObjectAdapter].set(() => arg) diff --git a/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFs.scala b/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFs.scala index e5db4d9c7384a..7472341eabb21 100644 --- a/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFs.scala +++ b/sql/hive/src/main/scala/org/apache/spark/sql/hive/hiveUDFs.scala @@ -25,7 +25,9 @@ import scala.jdk.CollectionConverters._ import org.apache.hadoop.hive.ql.exec._ import org.apache.hadoop.hive.ql.udf.generic._ import org.apache.hadoop.hive.ql.udf.generic.GenericUDAFEvaluator.AggregationBuffer -import org.apache.hadoop.hive.serde2.objectinspector.{ConstantObjectInspector, ObjectInspector, ObjectInspectorFactory} +import org.apache.hadoop.hive.serde2.objectinspector.{ConstantObjectInspector, ObjectInspector} +import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspectorFactory +import org.apache.hadoop.hive.serde2.objectinspector.StructObjectInspector import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ @@ -113,8 +115,17 @@ private[hive] case class HiveSimpleUDF( } } +/** + * A Hive GenericUDF whose Catalyst return type is snapshotted at analysis. + * + * Example: `SELECT hive_upper(CAST('Ab' AS CHAR(5)))` stores `CharType(5)` on `dataType` so + * executor conversion still pads even if first-class types are disabled there. + */ private[hive] case class HiveGenericUDF( - name: String, funcWrapper: HiveFunctionWrapper, children: Seq[Expression]) + name: String, + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + override val dataType: DataType) extends Expression with HiveInspectors with UserDefinedExpression { @@ -130,10 +141,9 @@ private[hive] case class HiveGenericUDF( override def foldable: Boolean = evaluator.isUDFDeterministic && evaluator.returnInspector.isInstanceOf[ConstantObjectInspector] - override lazy val dataType: DataType = inspectorToDataType(evaluator.returnInspector) - @transient - private lazy val evaluator = new HiveGenericUDFEvaluator(funcWrapper, children) + private lazy val evaluator = + new HiveGenericUDFEvaluator(funcWrapper, children, dataType) override def eval(input: InternalRow): Any = { children.zipWithIndex.foreach { @@ -193,6 +203,19 @@ private[hive] case class HiveGenericUDF( } } +object HiveGenericUDF { + def apply( + name: String, + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression]): HiveGenericUDF = { + HiveGenericUDF( + name, + funcWrapper, + children, + HiveGenericUDFEvaluator.inferReturnType(funcWrapper, children)) + } +} + /** * Converts a Hive Generic User Defined Table Generating Function (UDTF) to a * `Generator`. Note that the semantics of Generators do not allow @@ -207,36 +230,25 @@ private[hive] case class HiveGenericUDF( private[hive] case class HiveGenericUDTF( name: String, funcWrapper: HiveFunctionWrapper, - children: Seq[Expression]) + children: Seq[Expression], + override val elementSchema: StructType) extends Generator with HiveInspectors with CodegenFallback with UserDefinedExpression { @transient - protected lazy val function: GenericUDTF = { - val fun: GenericUDTF = funcWrapper.createFunction() - fun.setCollector(collector) - fun - } + protected lazy val collector = new UDTFCollector @transient - protected lazy val inputInspector = { - val inspectors = children.map(toInspector) - val fields = inspectors.indices.map(index => s"_col$index").asJava - ObjectInspectorFactory.getStandardStructObjectInspector(fields, inspectors.asJava) - } + private lazy val initialized = HiveGenericUDTF.initialize( + funcWrapper, children, collector, Some(elementSchema)) @transient - protected lazy val outputInspector = function.initialize(inputInspector) + protected lazy val function: GenericUDTF = initialized.function @transient - protected lazy val udtInput = new Array[AnyRef](children.length) + protected lazy val outputInspector = initialized.outputInspector @transient - protected lazy val collector = new UDTFCollector - - override lazy val elementSchema = StructType(outputInspector.getAllStructFieldRefs.asScala.map { - field => StructField(field.getFieldName, inspectorToDataType(field.getFieldObjectInspector), - nullable = true) - }.toArray) + protected lazy val udtInput = new Array[AnyRef](children.length) @transient private lazy val inputDataTypes: Array[DataType] = children.map(_.dataType).toArray @@ -245,7 +257,7 @@ private[hive] case class HiveGenericUDTF( private lazy val wrappers = children.map(x => wrapperFor(toInspector(x), x.dataType)).toArray @transient - private lazy val unwrapper = unwrapperFor(outputInspector) + private lazy val unwrapper = unwrapperFor(outputInspector, elementSchema) @transient private lazy val inputProjection = new InterpretedProjection(children) @@ -289,6 +301,52 @@ private[hive] case class HiveGenericUDTF( copy(children = newChildren) } +object HiveGenericUDTF extends HiveInspectors { + private[hive] case class InitializedUDTF( + function: GenericUDTF, + outputInspector: StructObjectInspector) + + def apply( + name: String, + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression]): HiveGenericUDTF = { + HiveGenericUDTF(name, funcWrapper, children, inferElementSchema(funcWrapper, children)) + } + + private[hive] def initialize( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + collector: Collector, + expectedSchema: Option[StructType] = None): InitializedUDTF = { + val function: GenericUDTF = funcWrapper.createFunction() + function.setCollector(collector) + val inspectors = children.map(toInspector) + val fields = inspectors.indices.map(index => s"_col$index").asJava + val inputInspector = + ObjectInspectorFactory.getStandardStructObjectInspector(fields, inspectors.asJava) + val outputInspector = function.initialize(inputInspector) + expectedSchema.foreach(checkCompatibleHiveReturnType(outputInspector, _)) + InitializedUDTF(function, outputInspector) + } + + def inferElementSchema( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression]): StructType = { + val initialized = initialize( + funcWrapper, + children, + new Collector { + override def collect(input: java.lang.Object): Unit = {} + }) + StructType(initialized.outputInspector.getAllStructFieldRefs.asScala.map { field => + StructField( + field.getFieldName, + inspectorToDataType(field.getFieldObjectInspector), + nullable = true) + }.toArray) + } +} + /** * While being evaluated by Spark SQL, the aggregation state of a Hive UDAF may be in the following * three formats: @@ -331,9 +389,11 @@ private[hive] case class HiveUDAFFunction( name: String, funcWrapper: HiveFunctionWrapper, children: Seq[Expression], - isUDAFBridgeRequired: Boolean = false, - mutableAggBufferOffset: Int = 0, - inputAggBufferOffset: Int = 0) + isUDAFBridgeRequired: Boolean, + mutableAggBufferOffset: Int, + inputAggBufferOffset: Int, + partialResultDataType: DataType, + override val dataType: DataType) extends TypedImperativeAggregate[HiveUDAFBuffer] with HiveInspectors with UserDefinedExpression { @@ -346,25 +406,18 @@ private[hive] case class HiveUDAFFunction( override def withNewInputAggBufferOffset(newInputAggBufferOffset: Int): ImperativeAggregate = copy(inputAggBufferOffset = newInputAggBufferOffset) - // Hive `ObjectInspector`s for all child expressions (input parameters of the function). @transient - private lazy val inputInspectors = children.map(toInspector).toArray + private lazy val initialized = HiveUDAFFunction.initializeEvaluators( + funcWrapper, + children, + isUDAFBridgeRequired, + Some(partialResultDataType), + Some(dataType)) // Spark SQL data types of input parameters. @transient private lazy val inputDataTypes: Array[DataType] = children.map(_.dataType).toArray - private def newEvaluator(): GenericUDAFEvaluator = { - val resolver = if (isUDAFBridgeRequired) { - new SparkGenericUDAFBridge(funcWrapper.createFunction[UDAF]()) - } else { - funcWrapper.createFunction[AbstractGenericUDAFResolver]() - } - - val parameterInfo = new SimpleGenericUDAFParameterInfo(inputInspectors, false, false, false) - resolver.getEvaluator(parameterInfo) - } - private case class HiveEvaluator( evaluator: GenericUDAFEvaluator, objectInspector: ObjectInspector) @@ -372,34 +425,20 @@ private[hive] case class HiveUDAFFunction( // The UDAF evaluator used to consume raw input rows and produce partial aggregation results. // Hive `ObjectInspector` used to inspect partial aggregation results. @transient - private lazy val partial1HiveEvaluator = { - val evaluator = newEvaluator() - HiveEvaluator(evaluator, evaluator.init(GenericUDAFEvaluator.Mode.PARTIAL1, inputInspectors)) - } + private lazy val partial1HiveEvaluator = HiveEvaluator( + initialized.partialEvaluator, initialized.partialInspector) // The UDAF evaluator used to consume partial aggregation results and produce final results. // Hive `ObjectInspector` used to inspect final results. @transient - private lazy val finalHiveEvaluator = { - val evaluator = newEvaluator() - HiveEvaluator( - evaluator, - evaluator.init(GenericUDAFEvaluator.Mode.FINAL, Array(partial1HiveEvaluator.objectInspector))) - } - - // Spark SQL data type of partial aggregation results - @transient - private lazy val partialResultDataType = - inspectorToDataType(partial1HiveEvaluator.objectInspector) + private lazy val finalHiveEvaluator = HiveEvaluator( + initialized.finalEvaluator, initialized.finalInspector) - // Wrapper functions used to wrap Spark SQL input arguments into Hive specific format. @transient private lazy val inputWrappers = children.map(x => wrapperFor(toInspector(x), x.dataType)).toArray - // Unwrapper function used to unwrap final aggregation result objects returned by Hive UDAFs into - // Spark SQL specific format. @transient - private lazy val resultUnwrapper = unwrapperFor(finalHiveEvaluator.objectInspector) + private lazy val resultUnwrapper = unwrapperFor(finalHiveEvaluator.objectInspector, dataType) @transient private lazy val cached: Array[AnyRef] = new Array[AnyRef](children.length) @@ -409,8 +448,6 @@ private[hive] case class HiveUDAFFunction( override def nullable: Boolean = true - override lazy val dataType: DataType = inspectorToDataType(finalHiveEvaluator.objectInspector) - override def prettyName: String = name override def sql(isDistinct: Boolean): String = { @@ -507,7 +544,8 @@ private[hive] case class HiveUDAFFunction( // Helper class used to de/serialize Hive UDAF `AggregationBuffer` objects private class AggregationBufferSerDe { - private val partialResultUnwrapper = unwrapperFor(partial1HiveEvaluator.objectInspector) + private val partialResultUnwrapper = + unwrapperFor(partial1HiveEvaluator.objectInspector, partialResultDataType) private val partialResultWrapper = wrapperFor(partial1HiveEvaluator.objectInspector, partialResultDataType) @@ -554,4 +592,77 @@ private[hive] case class HiveUDAFFunction( copy(children = newChildren) } +object HiveUDAFFunction extends HiveInspectors { + private[hive] case class InitializedEvaluators( + partialEvaluator: GenericUDAFEvaluator, + partialInspector: ObjectInspector, + finalEvaluator: GenericUDAFEvaluator, + finalInspector: ObjectInspector) + + def apply( + name: String, + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression]): HiveUDAFFunction = { + apply(name, funcWrapper, children, isUDAFBridgeRequired = false) + } + + def apply( + name: String, + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + isUDAFBridgeRequired: Boolean): HiveUDAFFunction = { + val (partialType, resultType) = + inferResolvedTypes(funcWrapper, children, isUDAFBridgeRequired) + HiveUDAFFunction( + name, + funcWrapper, + children, + isUDAFBridgeRequired, + mutableAggBufferOffset = 0, + inputAggBufferOffset = 0, + partialType, + resultType) + } + + private[hive] def initializeEvaluators( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + isUDAFBridgeRequired: Boolean, + expectedPartialType: Option[DataType] = None, + expectedResultType: Option[DataType] = None): InitializedEvaluators = { + val inputInspectors = children.map(toInspector).toArray + def newEvaluator(): GenericUDAFEvaluator = { + val resolver = if (isUDAFBridgeRequired) { + new SparkGenericUDAFBridge(funcWrapper.createFunction[UDAF]()) + } else { + funcWrapper.createFunction[AbstractGenericUDAFResolver]() + } + val parameterInfo = new SimpleGenericUDAFParameterInfo( + inputInspectors, false, false, false) + resolver.getEvaluator(parameterInfo) + } + val partial1 = newEvaluator() + val partialInspector = partial1.init(GenericUDAFEvaluator.Mode.PARTIAL1, inputInspectors) + val finalEvaluator = newEvaluator() + val finalInspector = + finalEvaluator.init(GenericUDAFEvaluator.Mode.FINAL, Array(partialInspector)) + expectedPartialType.foreach(checkCompatibleHiveReturnType(partialInspector, _)) + expectedResultType.foreach(checkCompatibleHiveReturnType(finalInspector, _)) + InitializedEvaluators( + partial1, + partialInspector, + finalEvaluator, + finalInspector) + } + + def inferResolvedTypes( + funcWrapper: HiveFunctionWrapper, + children: Seq[Expression], + isUDAFBridgeRequired: Boolean): (DataType, DataType) = { + val initialized = initializeEvaluators(funcWrapper, children, isUDAFBridgeRequired) + (inspectorToDataType(initialized.partialInspector), + inspectorToDataType(initialized.finalInspector)) + } +} + case class HiveUDAFBuffer(buf: AggregationBuffer, canDoMerge: Boolean) diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveInspectorSuite.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveInspectorSuite.scala index a88dcd0ef56aa..9f135bf6bc731 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveInspectorSuite.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveInspectorSuite.scala @@ -20,22 +20,30 @@ package org.apache.spark.sql.hive import java.util import org.apache.hadoop.hive.ql.udf.UDAFPercentile -import org.apache.hadoop.hive.serde2.io.DoubleWritable -import org.apache.hadoop.hive.serde2.objectinspector.{ObjectInspector, ObjectInspectorFactory, PrimitiveObjectInspector, StructObjectInspector} +import org.apache.hadoop.hive.serde2.io.{DoubleWritable, HiveCharWritable, HiveVarcharWritable} +import org.apache.hadoop.hive.serde2.objectinspector.{ConstantObjectInspector, ObjectInspector, ObjectInspectorFactory, PrimitiveObjectInspector, StructObjectInspector} import org.apache.hadoop.hive.serde2.objectinspector.ObjectInspectorFactory.ObjectInspectorOptions import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory -import org.apache.hadoop.hive.serde2.typeinfo.DecimalTypeInfo -import org.apache.hadoop.io.LongWritable +import org.apache.hadoop.hive.serde2.typeinfo.{CharTypeInfo, DecimalTypeInfo, VarcharTypeInfo} +import org.apache.hadoop.io.{LongWritable, Text} -import org.apache.spark.SparkFunSuite +import org.apache.spark.{SparkException, SparkFunSuite, SparkRuntimeException} import org.apache.spark.sql.{AnalysisException, Row, TestUserClassUDT} import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.Literal import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, GenericArrayData, MapData} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String class HiveInspectorSuite extends SparkFunSuite with HiveInspectors { + private def withFirstClassCharVarchar(enabled: Boolean)(f: => Unit): Unit = { + val conf = new SQLConf + conf.setConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS, enabled) + SQLConf.withExistingConf(conf)(f) + } + def unwrap(data: Any, oi: ObjectInspector): Any = { val unwrapper = unwrapperFor(oi) unwrapper(data) @@ -292,6 +300,283 @@ class HiveInspectorSuite extends SparkFunSuite with HiveInspectors { assert(typeInfo2.scale() === 10) } + test("SPARK-59277: Hive object inspectors preserve CHAR/VARCHAR type information") { + withFirstClassCharVarchar(enabled = true) { + Seq[DataType](CharType(5), VarcharType(7)).foreach { dataType => + val inspector = toInspector(dataType).asInstanceOf[PrimitiveObjectInspector] + assert(inspectorToDataType(inspector) === dataType) + dataType match { + case c: CharType => + assert(inspector.getTypeInfo.asInstanceOf[CharTypeInfo].getLength === c.length) + case v: VarcharType => + assert(inspector.getTypeInfo.asInstanceOf[VarcharTypeInfo].getLength === v.length) + } + } + } + } + + test("SPARK-59277: Hive object inspectors accept collated CHAR/VARCHAR values") { + withFirstClassCharVarchar(enabled = true) { + Seq[DataType]( + CharType(5, "UTF8_LCASE"), + VarcharType(7, "UNICODE_CI")).foreach { dataType => + val inspector = toInspector(dataType) + val value = UTF8String.fromString(dataType match { + case _: CharType => "ab" + case _: VarcharType => "abc" + }) + val expectedValue = dataType match { + case _: CharType => UTF8String.fromString("ab ") + case _: VarcharType => value + } + val expectedType = dataType match { + case c: CharType => CharType(c.length) + case v: VarcharType => VarcharType(v.length) + } + assert(inspectorToDataType(inspector) === expectedType) + assert(unwrap(wrap(value, inspector, dataType), inspector) === expectedValue) + } + } + } + + test("SPARK-59277: writable CHAR/VARCHAR inspectors round-trip values and nulls") { + withFirstClassCharVarchar(enabled = true) { + val charType = CharType(5) + val charInspector = PrimitiveObjectInspectorFactory.getPrimitiveWritableObjectInspector( + new CharTypeInfo(charType.length)) + val charValue = UTF8String.fromString("abc") + val wrappedChar = wrap(charValue, charInspector, charType) + assert(wrappedChar.isInstanceOf[HiveCharWritable]) + assert(wrappedChar.asInstanceOf[HiveCharWritable].getHiveChar.getPaddedValue === "abc ") + assert(unwrapperFor(charInspector, charType)(wrappedChar) === + UTF8String.fromString("abc ")) + assert(wrap(null, charInspector, charType) === null) + assert(unwrapperFor(charInspector, charType)(null) === null) + checkError( + exception = intercept[SparkRuntimeException] { + wrap(UTF8String.fromString("abcdef"), charInspector, charType) + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "5")) + + val varcharType = VarcharType(7) + val varcharInspector = PrimitiveObjectInspectorFactory.getPrimitiveWritableObjectInspector( + new VarcharTypeInfo(varcharType.length)) + val varcharValue = UTF8String.fromString("abc") + val wrappedVarchar = wrap(varcharValue, varcharInspector, varcharType) + assert(wrappedVarchar.isInstanceOf[HiveVarcharWritable]) + assert(wrappedVarchar.asInstanceOf[HiveVarcharWritable].getHiveVarchar.getValue === "abc") + assert(unwrapperFor(varcharInspector, varcharType)(wrappedVarchar) === varcharValue) + assert(wrap(null, varcharInspector, varcharType) === null) + assert(unwrapperFor(varcharInspector, varcharType)(null) === null) + checkError( + exception = intercept[SparkRuntimeException] { + wrap(UTF8String.fromString("abcdefgh"), varcharInspector, varcharType) + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "7")) + } + } + + test("SPARK-59277: Hive object inspectors support nested CHAR/VARCHAR") { + withFirstClassCharVarchar(enabled = true) { + val dataType = StructType(Seq( + StructField("chars", ArrayType(CharType(4))), + StructField("varchars", MapType(IntegerType, VarcharType(8))))) + val inspector = toInspector(dataType) + assert(inspectorToDataType(inspector) === dataType) + + val input = InternalRow( + new GenericArrayData(Array[Any](UTF8String.fromString("a"))), + ArrayBasedMapData( + Array[Any](1), + Array[Any](UTF8String.fromString("value")))) + val result = unwrapperFor(inspector, dataType)( + wrap(input, inspector, dataType)).asInstanceOf[InternalRow] + assert(result.getArray(0).getUTF8String(0) === UTF8String.fromString("a ")) + assert(result.getMap(1).valueArray().getUTF8String(0) === UTF8String.fromString("value")) + } + } + + test("SPARK-59277: Hive constant inspectors preserve CHAR/VARCHAR type information") { + withFirstClassCharVarchar(enabled = true) { + Seq[DataType](CharType(5), VarcharType(7)).foreach { dataType => + val value = UTF8String.fromString("abc") + val inspector = toInspector(Literal.create(value, dataType)) + assert(inspector.isInstanceOf[ConstantObjectInspector]) + assert(inspectorToDataType(inspector) === dataType) + val expected = dataType match { + case _: CharType => UTF8String.fromString("abc ") + case _: VarcharType => value + } + assert(unwrapperFor(inspector, dataType)( + inspector.asInstanceOf[ConstantObjectInspector].getWritableConstantValue) === expected) + + val nullInspector = toInspector(Literal.create(null, dataType)) + assert(nullInspector.isInstanceOf[ConstantObjectInspector]) + assert(inspectorToDataType(nullInspector) === dataType) + assert(unwrapperFor(nullInspector, dataType)( + nullInspector.asInstanceOf[ConstantObjectInspector].getWritableConstantValue) === null) + } + } + } + + test("SPARK-59277: Hive CHAR/VARCHAR inspectors remain STRING under legacy semantics") { + withFirstClassCharVarchar(enabled = false) { + Seq[DataType](CharType(5), VarcharType(7)).foreach { dataType => + val inspector = toInspector(dataType) + assert(inspectorToDataType(inspector) === StringType) + assert(inspectorToDataType(inspector, preserveCharVarchar = true) === dataType) + } + } + } + + test("SPARK-59277: Hive return type compatibility allows only STRING boundary drift") { + checkCompatibleHiveReturnType(StringType, CharType(5)) + checkCompatibleHiveReturnType(CharType(5), StringType) + checkCompatibleHiveReturnType(ArrayType(StringType), ArrayType(VarcharType(7))) + checkCompatibleHiveReturnType(CharType(5), CharType(5)) + checkCompatibleHiveReturnType(VarcharType(3), VarcharType(3)) + checkCompatibleHiveReturnType( + MapType(StringType, VarcharType(3)), + MapType(CharType(5), StringType)) + checkCompatibleHiveReturnType( + StructType.fromDDL("c STRING, v VARCHAR(3)"), + StructType.fromDDL("c CHAR(5), v STRING")) + + Seq[(DataType, DataType)]( + VarcharType(3) -> CharType(5), + CharType(4) -> CharType(5), + VarcharType(4) -> VarcharType(5), + IntegerType -> CharType(5), + MapType(CharType(4), StringType) -> MapType(CharType(5), StringType), + StructType.fromDDL("c STRING") -> StructType.fromDDL("c STRING, v STRING"), + StructType.fromDDL("c CHAR(4)") -> StructType.fromDDL("c CHAR(5)"), + StructType.fromDDL("a INT, b INT") -> StructType.fromDDL("b INT, a INT")).foreach { + case (runtimeType, expectedType) => + intercept[SparkException] { + checkCompatibleHiveReturnType(runtimeType, expectedType) + } + } + } + + test("SPARK-59277: Hive CHAR/VARCHAR boundaries enforce Spark length semantics") { + withFirstClassCharVarchar(enabled = true) { + val varchar = VarcharType(3) + val varcharInspector = toInspector(varchar) + checkError( + exception = intercept[SparkRuntimeException] { + wrap(UTF8String.fromString("abcd"), varcharInspector, varchar) + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "3")) + + val char = CharType(5) + val charInspector = toInspector(char) + checkError( + exception = intercept[SparkRuntimeException] { + wrap(UTF8String.fromString("abcdef"), charInspector, char) + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "5")) + + val varcharUnwrapper = + unwrapperFor(PrimitiveObjectInspectorFactory.javaStringObjectInspector, varchar) + checkError( + exception = intercept[SparkRuntimeException] { + varcharUnwrapper("abcd") + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "3")) + + val charUnwrapper = + unwrapperFor(PrimitiveObjectInspectorFactory.javaStringObjectInspector, char) + assert(charUnwrapper("ab") === UTF8String.fromString("ab ")) + checkError( + exception = intercept[SparkRuntimeException] { + charUnwrapper("abcdef") + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "5")) + } + } + + test("SPARK-59277: Hive map unwrapper validates CHAR-padded key collisions") { + withFirstClassCharVarchar(enabled = true) { + val targetType = MapType(CharType(2), StringType) + + // Two distinct keys that both pad to the same CHAR(2). Use plain-string + // inspectors so the Java HashMap preserves both entries (HiveChar + // equality would deduplicate them before our code runs). + val javaKeyOI = PrimitiveObjectInspectorFactory.javaStringObjectInspector + val javaValueOI = PrimitiveObjectInspectorFactory.javaStringObjectInspector + val javaMapOI = ObjectInspectorFactory + .getStandardMapObjectInspector(javaKeyOI, javaValueOI) + val javaMap = new java.util.LinkedHashMap[Any, Any]() + javaMap.put("a", "v1") + javaMap.put("a ", "v2") + + val writableKeyOI = + PrimitiveObjectInspectorFactory.writableStringObjectInspector + val writableValueOI = + PrimitiveObjectInspectorFactory.writableStringObjectInspector + val writableMapOI = ObjectInspectorFactory + .getStandardMapObjectInspector(writableKeyOI, writableValueOI) + val writableMap = new java.util.LinkedHashMap[Any, Any]() + writableMap.put(new Text("a"), new Text("w1")) + writableMap.put(new Text("a "), new Text("w2")) + + Seq( + ("java", javaMapOI, javaMap, "v2"), + ("writable", writableMapOI, writableMap, "w2") + ).foreach { case (label, mapOI, map, lastValue) => + // EXCEPTION policy: duplicate padded keys must throw. + val exConf = new SQLConf + exConf.setConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS, true) + exConf.setConfString( + SQLConf.MAP_KEY_DEDUP_POLICY.key, "EXCEPTION") + SQLConf.withExistingConf(exConf) { + val unwrapper = unwrapperFor(mapOI, targetType) + intercept[SparkRuntimeException] { unwrapper(map) } + } + + // LAST_WIN policy: last inserted value wins. + val lwConf = new SQLConf + lwConf.setConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS, true) + lwConf.setConfString( + SQLConf.MAP_KEY_DEDUP_POLICY.key, "LAST_WIN") + SQLConf.withExistingConf(lwConf) { + val unwrapper = unwrapperFor(mapOI, targetType) + val result = unwrapper(map).asInstanceOf[MapData] + assert(result.numElements() === 1, + s"$label: expected 1 entry after dedup") + assert(result.keyArray().getUTF8String(0) === + UTF8String.fromString("a ")) + assert(result.valueArray().getUTF8String(0) === + UTF8String.fromString(lastValue), + s"$label: last-win value") + } + } + } + } + + test("SPARK-59277: Hive object inspectors reject unsupported CHAR/VARCHAR lengths") { + withFirstClassCharVarchar(enabled = true) { + Seq[DataType]( + CharType(0), CharType(256), VarcharType(0), VarcharType(65536)).foreach { dataType => + val expectedParams = Map("typeName" -> s"\"${dataType.sql}\"") + checkError( + exception = intercept[AnalysisException](toInspector(dataType)), + condition = "UNSUPPORTED_DATATYPE", + parameters = expectedParams) + checkError( + exception = intercept[AnalysisException](dataType.toTypeInfo), + condition = "UNSUPPORTED_DATATYPE", + parameters = expectedParams) + } + } + } + test("SPARK-57556: TIME type is unsupported in Hive object inspectors") { val timeType = TimeType() val expectedParams = Map("typeName" -> s"\"${timeType.sql}\"") diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveScalaReflectionSuite.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveScalaReflectionSuite.scala index ce46baae9e468..b777bacfe4fee 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveScalaReflectionSuite.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/HiveScalaReflectionSuite.scala @@ -29,7 +29,7 @@ class HiveScalaReflectionSuite extends SparkFunSuite { test("SPARK-38510: ScalaReflection.getConstructorParameterNames should work for classes with " + "cyclic annotation references") { - assert(Seq("name", "funcWrapper", "children") === + assert(Seq("name", "funcWrapper", "children", "dataType") === ScalaReflection.getConstructorParameterNames(classOf[HiveGenericUDF])) } } diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDAFSuite.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDAFSuite.scala index 9fec3e58de255..b266a5feba384 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDAFSuite.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDAFSuite.scala @@ -19,24 +19,30 @@ package org.apache.spark.sql.hive.execution import scala.jdk.CollectionConverters._ +import org.apache.hadoop.hive.common.`type`.HiveChar import org.apache.hadoop.hive.ql.udf.UDAFPercentile import org.apache.hadoop.hive.ql.udf.generic.{AbstractGenericUDAFResolver, GenericUDAFEvaluator, GenericUDAFMax} import org.apache.hadoop.hive.ql.udf.generic.GenericUDAFEvaluator.{AggregationBuffer, Mode} import org.apache.hadoop.hive.ql.util.JavaDataModel -import org.apache.hadoop.hive.serde2.objectinspector.{ObjectInspector, ObjectInspectorFactory} -import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory -import org.apache.hadoop.hive.serde2.typeinfo.TypeInfo +import org.apache.hadoop.hive.serde2.objectinspector.{ObjectInspector, ObjectInspectorFactory, PrimitiveObjectInspector} +import org.apache.hadoop.hive.serde2.objectinspector.primitive.{PrimitiveObjectInspectorFactory, PrimitiveObjectInspectorUtils} +import org.apache.hadoop.hive.serde2.typeinfo.{CharTypeInfo, TypeInfo} import test.org.apache.spark.sql.MyDoubleAvg -import org.apache.spark.SPARK_DOC_ROOT +import org.apache.spark.{SPARK_DOC_ROOT, SparkException} import org.apache.spark.sql.{AnalysisException, DataFrame, QueryTest, Row} +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, Literal} import org.apache.spark.sql.catalyst.expressions.Cast._ import org.apache.spark.sql.catalyst.expressions.aggregate.Complete import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.execution.aggregate.ObjectHashAggregateExec +import org.apache.spark.sql.hive.HiveShim.HiveFunctionWrapper +import org.apache.spark.sql.hive.HiveUDAFFunction import org.apache.spark.sql.hive.test.TestHiveSingleton import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{CharType, LongType, StringType, VarcharType} import org.apache.spark.tags.SlowHiveTest +import org.apache.spark.unsafe.types.UTF8String @SlowHiveTest class HiveUDAFSuite extends QueryTest @@ -200,6 +206,90 @@ class HiveUDAFSuite extends QueryTest } } + test("SPARK-59277: Hive UDAF supports first-class CHAR/VARCHAR") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + Seq( + ("CHAR(5) COLLATE UTF8_LCASE", CharType(5), Row("def ")), + ("VARCHAR(7) COLLATE UNICODE_CI", VarcharType(7), Row("def")) + ).foreach { case (dataType, expectedType, expectedRow) => + val aggregate = sql( + s"""SELECT hive_max(value) + |FROM VALUES + | (CAST('abc' AS $dataType)), + | (CAST('def' AS $dataType)) + |AS input(value) + |""".stripMargin) + assert(aggregate.schema.head.dataType === expectedType) + checkAnswer(aggregate, expectedRow) + } + } + } + + test("SPARK-59277: HiveUDAFFunction keeps analysis type after child constantness changes") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val attr = AttributeReference("value", VarcharType(7))() + val original = HiveUDAFFunction( + "hive_max", + HiveFunctionWrapper(classOf[GenericUDAFMax].getName), + Seq(Literal.create(UTF8String.fromString("abc"), VarcharType(7)))) + assert(original.dataType === VarcharType(7)) + val copied = original.withNewChildren(Seq(attr)).asInstanceOf[HiveUDAFFunction] + assert(copied.dataType === original.dataType) + copied.serialize(null) + } + } + + test("SPARK-59277: Hive UDAF partial buffer type can differ from the final result") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + withUserDefinedFunction("char_max" -> true) { + sql( + s"CREATE TEMPORARY FUNCTION char_max AS " + + s"'${classOf[MockPartialStringFinalCharUDAF].getName}'") + withTempView("cv_udaf") { + Seq("abc", "def").toDF("value").repartition(2).createOrReplaceTempView("cv_udaf") + val aggregate = sql("SELECT char_max(CAST(value AS VARCHAR(7))) FROM cv_udaf") + assert(aggregate.schema.head.dataType === CharType(5)) + checkAnswer(aggregate, Row("def ")) + } + } + } + } + + test("SPARK-59277: incompatible UDAF partial inspector triggers mismatch error") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + // MockPartialStringFinalCharUDAF exposes STRING partial / CHAR(5) final. + // Feed LongType as the expected partial type to trigger the mismatch. + val udaf = HiveUDAFFunction( + "char_max", + HiveFunctionWrapper(classOf[MockPartialStringFinalCharUDAF].getName), + Seq(Literal("x")), + isUDAFBridgeRequired = false, + mutableAggBufferOffset = 0, + inputAggBufferOffset = 0, + partialResultDataType = LongType, + dataType = CharType(5)) + intercept[SparkException] { udaf.serialize(null) } + } + } + + test("SPARK-59277: incompatible UDAF final inspector triggers mismatch error") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + // MockPartialStringFinalCharUDAF exposes CHAR(5) final. + // Feed VarcharType(5) as the expected final type: CHAR(5) vs VARCHAR(5) + // is incompatible (different bounded-string kind). + val udaf = HiveUDAFFunction( + "char_max", + HiveFunctionWrapper(classOf[MockPartialStringFinalCharUDAF].getName), + Seq(Literal("x")), + isUDAFBridgeRequired = false, + mutableAggBufferOffset = 0, + inputAggBufferOffset = 0, + partialResultDataType = StringType, + dataType = VarcharType(5)) + intercept[SparkException] { udaf.serialize(null) } + } + } + test("non-deterministic children expressions of UDAF") { withTempView("view1") { spark.range(1).selectExpr("id as x", "id as y").createTempView("view1") @@ -405,3 +495,67 @@ class MockUDAFEvaluator2 extends GenericUDAFEvaluator { Array[Object](buffer.nonNullCount: java.lang.Long, buffer.nullCount: java.lang.Long) } } + +/** + * PARTIAL1/PARTIAL2 expose a STRING inspector; FINAL/COMPLETE expose CHAR(5). This keeps the + * (partial, final) Catalyst type pair distinct so shuffle serde cannot silently use the result + * type for the aggregation buffer. + */ +class MockPartialStringFinalCharUDAF extends AbstractGenericUDAFResolver { + override def getEvaluator(info: Array[TypeInfo]): GenericUDAFEvaluator = + new MockPartialStringFinalCharEvaluator +} + +class MockPartialStringFinalCharBuffer(var max: String) + extends GenericUDAFEvaluator.AbstractAggregationBuffer { + override def estimate(): Int = 16 +} + +class MockPartialStringFinalCharEvaluator extends GenericUDAFEvaluator { + private var inputOI: PrimitiveObjectInspector = _ + private val partialOI = PrimitiveObjectInspectorFactory.javaStringObjectInspector + private val finalOI = + PrimitiveObjectInspectorFactory.getPrimitiveJavaObjectInspector(new CharTypeInfo(5)) + + override def init(mode: Mode, parameters: Array[ObjectInspector]): ObjectInspector = { + if (mode == Mode.PARTIAL1 || mode == Mode.COMPLETE) { + inputOI = parameters.head.asInstanceOf[PrimitiveObjectInspector] + } + if (mode == Mode.PARTIAL1 || mode == Mode.PARTIAL2) partialOI else finalOI + } + + override def getNewAggregationBuffer: AggregationBuffer = + new MockPartialStringFinalCharBuffer(null) + + override def reset(agg: AggregationBuffer): Unit = { + agg.asInstanceOf[MockPartialStringFinalCharBuffer].max = null + } + + override def iterate(agg: AggregationBuffer, parameters: Array[AnyRef]): Unit = { + if (parameters.head != null) { + val value = PrimitiveObjectInspectorUtils.getString(parameters.head, inputOI) + val buffer = agg.asInstanceOf[MockPartialStringFinalCharBuffer] + if (buffer.max == null || value > buffer.max) { + buffer.max = value + } + } + } + + override def merge(agg: AggregationBuffer, partial: Object): Unit = { + if (partial != null) { + val value = partial.asInstanceOf[String] + val buffer = agg.asInstanceOf[MockPartialStringFinalCharBuffer] + if (buffer.max == null || value > buffer.max) { + buffer.max = value + } + } + } + + override def terminatePartial(agg: AggregationBuffer): AnyRef = + agg.asInstanceOf[MockPartialStringFinalCharBuffer].max + + override def terminate(agg: AggregationBuffer): AnyRef = { + val max = agg.asInstanceOf[MockPartialStringFinalCharBuffer].max + if (max == null) null else new HiveChar(max, 5) + } +} diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDFSuite.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDFSuite.scala index 942172d1411c3..efba2248819a2 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDFSuite.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/execution/HiveUDFSuite.scala @@ -17,13 +17,15 @@ package org.apache.spark.sql.hive.execution -import java.io.{DataInput, DataOutput, File, PrintWriter} +import java.io.{ByteArrayInputStream, ByteArrayOutputStream, DataInput, DataOutput, File} +import java.io.{ObjectInputStream, ObjectOutputStream, PrintWriter} import java.sql.{Date, Timestamp} import java.util.{ArrayList, Arrays, Properties} import scala.jdk.CollectionConverters._ import org.apache.hadoop.conf.Configuration +import org.apache.hadoop.hive.common.`type`.HiveChar import org.apache.hadoop.hive.ql.exec.UDF import org.apache.hadoop.hive.ql.metadata.HiveException import org.apache.hadoop.hive.ql.udf.{UDAFPercentile, UDFType} @@ -31,23 +33,27 @@ import org.apache.hadoop.hive.ql.udf.generic._ import org.apache.hadoop.hive.ql.udf.generic.GenericUDF.DeferredObject import org.apache.hadoop.hive.serde2.{AbstractSerDe, SerDeStats} import org.apache.hadoop.hive.serde2.objectinspector.{ObjectInspector, ObjectInspectorFactory} -import org.apache.hadoop.hive.serde2.objectinspector.primitive.PrimitiveObjectInspectorFactory +import org.apache.hadoop.hive.serde2.objectinspector.primitive.{ + HiveCharObjectInspector, + PrimitiveObjectInspectorFactory} +import org.apache.hadoop.hive.serde2.typeinfo.TypeInfoFactory import org.apache.hadoop.io.{LongWritable, Writable} -import org.apache.spark.{SparkException, SparkFiles, TestUtils} +import org.apache.spark.{SparkException, SparkFiles, SparkRuntimeException, TestUtils} import org.apache.spark.sql.{AnalysisException, QueryTest, Row} import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{AttributeReference, BindReferences, CodegenObjectFactoryMode, Literal} import org.apache.spark.sql.catalyst.plans.logical.{Filter, Project} -import org.apache.spark.sql.catalyst.util.DateTimeUtils +import org.apache.spark.sql.catalyst.util.{DateTimeUtils, GenericArrayData} import org.apache.spark.sql.execution.WholeStageCodegenExec import org.apache.spark.sql.functions.{call_function, max} -import org.apache.spark.sql.hive.HiveGenericUDF +import org.apache.spark.sql.hive.{HiveGenericUDF, HiveGenericUDTF} import org.apache.spark.sql.hive.HiveShim.HiveFunctionWrapper import org.apache.spark.sql.hive.test.{TestHiveSingleton, TestUDTFJar} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{TimestampType, TimeType} +import org.apache.spark.sql.types._ import org.apache.spark.tags.SlowHiveTest +import org.apache.spark.unsafe.types.UTF8String import org.apache.spark.util.Utils case class Fields(f1: Int, f2: Int, f3: Int, f4: Int, f5: Int) @@ -892,6 +898,265 @@ class HiveUDFSuite extends QueryTest with TestHiveSingleton { hiveContext.reset() } + test("SPARK-59277: Hive UDF and UDTF support first-class CHAR/VARCHAR") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + withUserDefinedFunction( + "hive_simple_concat" -> true, + "hive_char_padded" -> true, + "hive_upper" -> true, + "hive_explode" -> true) { + sql(s"CREATE TEMPORARY FUNCTION hive_simple_concat AS " + + s"'${classOf[UDFStringString].getName}'") + sql(s"CREATE TEMPORARY FUNCTION hive_char_padded AS " + + s"'${classOf[InspectCharGenericUDF].getName}'") + sql(s"CREATE TEMPORARY FUNCTION hive_upper AS '${classOf[GenericUDFUpper].getName}'") + sql(s"CREATE TEMPORARY FUNCTION hive_explode AS '${classOf[GenericUDTFExplode].getName}'") + + val simple = sql( + """SELECT hive_simple_concat( + | CAST('A' AS CHAR(3)), + | CAST('b' AS VARCHAR(2))) AS value""".stripMargin) + assert(simple.schema.head.dataType === StringType) + // Hive's simple-UDF conversion from CHAR to a Java String strips trailing spaces. + checkAnswer(simple, Row("A b")) + + // A GenericUDF reading the HiveChar directly observes the padded CHAR value. + checkAnswer( + sql("SELECT hive_char_padded(CAST('A' AS CHAR(3)))"), + Row("A ")) + + val scalar = sql( + "SELECT hive_upper(CAST('Ab' AS CHAR(5) COLLATE UTF8_LCASE)) AS value") + assert(scalar.schema.head.dataType === CharType(5)) + checkAnswer(scalar, Row("AB ")) + + val table = sql( + """SELECT value + |FROM ( + | SELECT array(CAST('abc' AS VARCHAR(7) COLLATE UNICODE_CI)) AS values + |) input + |LATERAL VIEW hive_explode(values) exploded AS value + |""".stripMargin) + assert(table.schema.head.dataType === VarcharType(7)) + checkAnswer(table, Row("abc")) + } + } + } + + test("SPARK-59277: Hive UDF supports preserved CHAR without standard semantics") { + withSQLConf( + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true", + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { + withUserDefinedFunction("hive_upper" -> true) { + sql(s"CREATE TEMPORARY FUNCTION hive_upper AS '${classOf[GenericUDFUpper].getName}'") + + val result = sql("SELECT hive_upper(CAST('Ab' AS CHAR(5))) AS value") + assert(result.schema.head.dataType === CharType(5)) + checkAnswer(result, Row("AB ")) + } + } + } + + test("SPARK-59277: persisted views use their captured Hive conversion types") { + withUserDefinedFunction( + "hive_upper" -> false, + "hive_max" -> false, + "hive_explode" -> false) { + withView("first_class_hive_view", "first_class_hive_udtf_view", "legacy_hive_view") { + sql(s"CREATE FUNCTION hive_upper AS '${classOf[GenericUDFUpper].getName}'") + sql(s"CREATE FUNCTION hive_max AS '${classOf[GenericUDAFMax].getName}'") + sql(s"CREATE FUNCTION hive_explode AS '${classOf[GenericUDTFExplode].getName}'") + val query = + """SELECT + | hive_upper(CAST('Ab' AS CHAR(5))) AS scalar_value, + | hive_max(CAST('cd' AS CHAR(4))) AS aggregate_value""".stripMargin + val udtfQuery = + """SELECT value + |FROM ( + | SELECT array(CAST('xy' AS CHAR(5))) AS values + |) input + |LATERAL VIEW hive_explode(values) exploded AS value + |""".stripMargin + + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + sql(s"CREATE VIEW first_class_hive_view AS $query") + sql(s"CREATE VIEW first_class_hive_udtf_view AS $udtfQuery") + } + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + SQLConf.CODEGEN_FACTORY_MODE.key -> CodegenObjectFactoryMode.NO_CODEGEN.toString) { + val result = sql("SELECT * FROM first_class_hive_view") + assert(result.schema.map(_.dataType) === Seq(StringType, StringType)) + checkAnswer(result, Row("AB ", "cd ")) + + val udtfResult = sql("SELECT * FROM first_class_hive_udtf_view") + assert(udtfResult.schema.head.dataType === StringType) + checkAnswer(udtfResult, Row("xy ")) + } + + withSQLConf( + SQLConf.LEGACY_CHAR_VARCHAR_AS_STRING.key -> "true", + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false") { + sql(s"CREATE VIEW legacy_hive_view AS $query") + } + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val result = sql("SELECT * FROM legacy_hive_view") + assert(result.schema.map(_.dataType) === Seq(StringType, StringType)) + checkAnswer(result, Row("AB", "cd")) + } + } + } + } + + test("SPARK-59277: HiveGenericUDF Java serialization preserves its Catalyst type") { + def serialize(expression: HiveGenericUDF): Array[Byte] = { + val bytes = new ByteArrayOutputStream() + val output = new ObjectOutputStream(bytes) + try { + output.writeObject(expression) + } finally { + output.close() + } + bytes.toByteArray + } + + def deserialize(bytes: Array[Byte]): HiveGenericUDF = { + val input = new ObjectInputStream(new ByteArrayInputStream(bytes)) + try { + input.readObject().asInstanceOf[HiveGenericUDF] + } finally { + input.close() + } + } + + val inferredBytes = withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + val expression = HiveGenericUDF( + "return_char", + HiveFunctionWrapper(classOf[ReturnCharGenericUDF].getName), + Seq(Literal("ab"))) + assert(expression.dataType === CharType(5)) + val copied = expression.withNewChildren(Seq(Literal("cd"))).asInstanceOf[HiveGenericUDF] + assert(copied.dataType === CharType(5)) + serialize(expression) + } + + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false") { + val expression = deserialize(inferredBytes) + assert(expression.dataType === CharType(5)) + assert(expression.eval(InternalRow.empty) === UTF8String.fromString("ab ")) + } + + val (charBytes, varcharBytes) = withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false") { + def resolvedExpression(value: String, dataType: DataType) = { + val expression = HiveGenericUDF( + "return_string", + HiveFunctionWrapper(classOf[ReturnStringGenericUDF].getName), + Seq(Literal(value)), + dataType) + assert(expression.dataType === dataType) + serialize(expression) + } + ( + resolvedExpression("ab", CharType(5)), + resolvedExpression("abcd", VarcharType(3))) + } + + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "true") { + val charExpression = deserialize(charBytes) + assert(charExpression.dataType === CharType(5)) + assert(charExpression.eval(InternalRow.empty) === UTF8String.fromString("ab ")) + + val varcharExpression = deserialize(varcharBytes) + assert(varcharExpression.dataType === VarcharType(3)) + val exception = intercept[SparkException] { + varcharExpression.eval(InternalRow.empty) + } + checkError( + exception = exception.getCause.asInstanceOf[SparkRuntimeException], + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "3")) + } + } + + test("SPARK-59277: HiveGenericUDF keeps analysis type after child constantness changes") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val attr = AttributeReference("value", CharType(5))() + val original = HiveGenericUDF( + "hive_upper", + HiveFunctionWrapper(classOf[GenericUDFUpper].getName), + Seq(Literal.create(UTF8String.fromString("Ab"), CharType(5)))) + assert(original.dataType === CharType(5)) + val copied = original.withNewChildren(Seq(attr)).asInstanceOf[HiveGenericUDF] + assert(copied.dataType === CharType(5)) + val bound = BindReferences.bindReference(copied, Seq(attr)) + assert(bound.eval(InternalRow(UTF8String.fromString("Ab "))) === + UTF8String.fromString("AB ")) + } + } + + test("SPARK-59277: HiveGenericUDTF keeps analysis schema after child constantness changes") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val arrayType = ArrayType(VarcharType(7)) + val values = new GenericArrayData(Array[Any](UTF8String.fromString("abc"))) + val original = HiveGenericUDTF( + "hive_explode", + HiveFunctionWrapper(classOf[GenericUDTFExplode].getName), + Seq(Literal(values, arrayType))) + assert(original.elementSchema.head.dataType === VarcharType(7)) + val attr = AttributeReference("values", arrayType)() + val copied = original.withNewChildren(Seq(attr)).asInstanceOf[HiveGenericUDTF] + assert(copied.elementSchema === original.elementSchema) + val bound = BindReferences.bindReference(copied, Seq(attr)) + .asInstanceOf[HiveGenericUDTF] + val rows = bound.eval(InternalRow(values)).iterator.toSeq + assert(rows.map(_.get(0, VarcharType(7))) === Seq(UTF8String.fromString("abc"))) + } + } + + test("SPARK-59277: incompatible UDF runtime inspector triggers mismatch error") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + // ReturnCharGenericUDF always returns CHAR(5), but we force VarcharType(5) as + // the captured analysis type. The runtime check must detect the mismatch. + val udf = HiveGenericUDF( + "return_char", + HiveFunctionWrapper(classOf[ReturnCharGenericUDF].getName), + Seq(Literal("x")), + VarcharType(5)) + intercept[SparkException] { udf.eval(InternalRow.empty) } + } + } + + test("SPARK-59277: incompatible UDTF runtime inspector triggers mismatch error") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + // GenericUDTFExplode returns the element type of its input. Feeding it + // ARRAY makes its runtime inspector return CHAR(5). We supply a + // captured schema claiming VARCHAR(5), so the check must fire. + val arrayType = ArrayType(CharType(5)) + val values = new GenericArrayData( + Array[Any](UTF8String.fromString("abc "))) + val schema = StructType(Seq(StructField("col", VarcharType(5)))) + val udtf = HiveGenericUDTF( + "hive_explode", + HiveFunctionWrapper(classOf[GenericUDTFExplode].getName), + Seq(Literal(values, arrayType)), + schema) + intercept[SparkException] { udtf.eval(InternalRow.empty) } + } + } + test("SPARK-58792: copied HiveGenericUDF nodes must not share a mutable GenericUDF") { val tsAttr = AttributeReference("ts", TimestampType, nullable = false)() val constTs = Literal( @@ -1025,6 +1290,41 @@ class PairUDF extends GenericUDF { override def getDisplayString(p1: Array[String]): String = "" } +class ReturnCharGenericUDF extends GenericUDF { + override def initialize(arguments: Array[ObjectInspector]): ObjectInspector = + PrimitiveObjectInspectorFactory.getPrimitiveJavaObjectInspector( + TypeInfoFactory.getCharTypeInfo(5)) + + override def evaluate(arguments: Array[DeferredObject]): AnyRef = + new HiveChar(arguments(0).get.toString, 5) + + override def getDisplayString(children: Array[String]): String = "return_char" +} + +class ReturnStringGenericUDF extends GenericUDF { + override def initialize(arguments: Array[ObjectInspector]): ObjectInspector = + PrimitiveObjectInspectorFactory.javaStringObjectInspector + + override def evaluate(arguments: Array[DeferredObject]): AnyRef = + arguments(0).get.toString + + override def getDisplayString(children: Array[String]): String = "return_string" +} + +class InspectCharGenericUDF extends GenericUDF { + private var inspector: HiveCharObjectInspector = _ + + override def initialize(arguments: Array[ObjectInspector]): ObjectInspector = { + inspector = arguments(0).asInstanceOf[HiveCharObjectInspector] + PrimitiveObjectInspectorFactory.javaStringObjectInspector + } + + override def evaluate(arguments: Array[DeferredObject]): AnyRef = + inspector.getPrimitiveJavaObject(arguments(0).get).getPaddedValue + + override def getDisplayString(children: Array[String]): String = "inspect_char" +} + @UDFType(stateful = true) class StatefulUDF extends UDF { private val result = new LongWritable(0)