diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/csv/UnivocityParser.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/csv/UnivocityParser.scala index 4583c328239ee..6f0965e2aed03 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/csv/UnivocityParser.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/csv/UnivocityParser.scala @@ -278,8 +278,10 @@ class UnivocityParser( timeFormatter.parse(datum) } - case _: StringType => (d: String) => - nullSafeDatum(d, name, nullable, options)(UTF8String.fromString) + case dt: StringType => (d: String) => + nullSafeDatum(d, name, nullable, options) { s => + CharVarcharUtils.applyTextParseSemantics(UTF8String.fromString(s), dt) + } case _: BinaryType => (d: String) => nullSafeDatum(d, name, nullable, options)(_.getBytes) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/csvExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/csvExpressions.scala index 15d4c15dcbdee..17035400aa6f1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/csvExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/csvExpressions.scala @@ -183,7 +183,7 @@ case class SchemaOfCsv( DataTypeMismatch( errorSubClass = "UNEXPECTED_NULL", messageParameters = Map("exprName" -> "csv")) - } else if (child.dataType != StringType) { + } else if (!child.dataType.isInstanceOf[StringType]) { DataTypeMismatch( errorSubClass = "UNEXPECTED_INPUT_TYPE", messageParameters = Map( diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/jsonExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/jsonExpressions.scala index 90867110c7407..47e21926880fc 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/jsonExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/jsonExpressions.scala @@ -1913,7 +1913,7 @@ case class SchemaOfJson( DataTypeMismatch( errorSubClass = "UNEXPECTED_NULL", messageParameters = Map("exprName" -> "json")) - } else if (child.dataType != StringType) { + } else if (!child.dataType.isInstanceOf[StringType]) { DataTypeMismatch( errorSubClass = "UNEXPECTED_INPUT_TYPE", messageParameters = Map( diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/xmlExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/xmlExpressions.scala index 31aea4910543b..cc5b989763e82 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/xmlExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/xmlExpressions.scala @@ -200,7 +200,7 @@ case class SchemaOfXml( DataTypeMismatch( errorSubClass = "UNEXPECTED_NULL", messageParameters = Map("exprName" -> "xml")) - } else if (child.dataType != StringType) { + } else if (!child.dataType.isInstanceOf[StringType]) { DataTypeMismatch( errorSubClass = "UNEXPECTED_INPUT_TYPE", messageParameters = Map( diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/json/JacksonParser.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/json/JacksonParser.scala index ea8061b774c43..1ddd6f8e95b10 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/json/JacksonParser.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/json/JacksonParser.scala @@ -27,7 +27,7 @@ import scala.util.control.NonFatal import com.fasterxml.jackson.core._ import org.apache.hadoop.fs.PositionedReadable -import org.apache.spark.SparkUpgradeException +import org.apache.spark.{SparkRuntimeException, SparkUpgradeException} import org.apache.spark.internal.Logging import org.apache.spark.sql.catalyst.{InternalRow, NoopFilters, StructFilters} import org.apache.spark.sql.catalyst.expressions._ @@ -40,7 +40,7 @@ import org.apache.spark.sql.sources.Filter import org.apache.spark.sql.types._ import org.apache.spark.types.variant._ import org.apache.spark.unsafe.types.{CalendarInterval, TimestampNanosVal, UTF8String, VariantVal} -import org.apache.spark.util.Utils +import org.apache.spark.util.{SparkErrorUtils, Utils} /** * Constructs a parser for a given schema that translates a json string to an [[InternalRow]]. @@ -54,6 +54,16 @@ class JacksonParser( import JacksonUtils._ import com.fasterxml.jackson.core.JsonToken._ + private object DuplicateMapKeyException { + def unapply(exception: Throwable): Option[SparkRuntimeException] = { + SparkErrorUtils.getRootCause(exception) match { + case cause: SparkRuntimeException if cause.getCondition == "DUPLICATED_MAP_KEY" => + Some(cause) + case _ => None + } + } + } + // A `ValueConverter` is responsible for converting a value from `JsonParser` // to a value in a field for `InternalRow`. private type ValueConverter = JsonParser => AnyRef @@ -187,7 +197,8 @@ class JacksonParser( private def makeMapRootConverter(mt: MapType): JsonParser => Iterable[InternalRow] = { val fieldConverter = makeConverter(mt.valueType) (parser: JsonParser) => parseJsonToken[Iterable[InternalRow]](parser, mt) { - case START_OBJECT => Some(InternalRow(convertMap(parser, fieldConverter))) + case START_OBJECT => + Some(InternalRow(convertMap(parser, fieldConverter, mt.keyType, mt.valueType))) } } @@ -302,7 +313,7 @@ class JacksonParser( } } - case _: StringType => (parser: JsonParser) => { + case dt: StringType => (parser: JsonParser) => { // This must be enabled if we will retrieve the bytes directly from the raw content: val oldFeature = parser.getFeatureMask val featureToAdd = JsonParser.Feature.INCLUDE_SOURCE_IN_LOCATION.getMask @@ -356,7 +367,7 @@ class JacksonParser( // to be reset. This ensures that every feature is restored to its previous // state as defined by `oldFeature`. parser.overrideStdFeatures(oldFeature, ~0) - result + CharVarcharUtils.applyTextParseSemantics(result, dt) } case TimestampType => @@ -480,7 +491,7 @@ class JacksonParser( case mt: MapType => val valueConverter = makeConverter(mt.valueType) (parser: JsonParser) => parseJsonToken[MapData](parser, dataType) { - case START_OBJECT => convertMap(parser, valueConverter) + case START_OBJECT => convertMap(parser, valueConverter, mt.keyType, mt.valueType) } case udt: UserDefinedType[_] => @@ -589,6 +600,7 @@ class JacksonParser( bitmask(index) = false } catch { case e: SparkUpgradeException => throw e + case DuplicateMapKeyException(e) => throw e case err: PartialValueException if enablePartialResults => badRecordException = badRecordException.orElse(Some(err.cause)) row.update(index, err.partialResult) @@ -617,33 +629,58 @@ class JacksonParser( */ private def convertMap( parser: JsonParser, - fieldConverter: ValueConverter): MapData = { + fieldConverter: ValueConverter, + keyType: DataType, + valueType: DataType): MapData = { val keys = ArrayBuffer.empty[UTF8String] val values = ArrayBuffer.empty[Any] - var badRecordException: Option[Throwable] = None + var partialResultException: Option[Throwable] = None + var badMapException: Option[Throwable] = None while (nextUntil(parser, JsonToken.END_OBJECT)) { - keys += UTF8String.fromString(parser.currentName) - try { - values += fieldConverter.apply(parser) + val rawKey = UTF8String.fromString(parser.currentName) + val value = try { + Some(fieldConverter.apply(parser)) } catch { case err: PartialValueException if enablePartialResults => - badRecordException = badRecordException.orElse(Some(err.cause)) - values += err.partialResult + partialResultException = partialResultException.orElse(Some(err.cause)) + Some(err.partialResult) + case DuplicateMapKeyException(e) => throw e case NonFatal(e) if enablePartialResults => - badRecordException = badRecordException.orElse(Some(e)) + badMapException = badMapException.orElse(Some(e)) parser.skipChildren() + None + } + value.foreach { parsedValue => + try { + val key = CharVarcharUtils.applyTextParseSemantics(rawKey, keyType) + keys += key + values += parsedValue + } catch { + case DuplicateMapKeyException(e) => throw e + case NonFatal(e) if enablePartialResults => + badMapException = badMapException.orElse(Some(e)) + } } } - // The JSON map will never have null or duplicated map keys, it's safe to create a - // ArrayBasedMapData directly here. - val mapData = ArrayBasedMapData(keys.toArray, values.toArray) + val mapData = keyType match { + case _: CharType | _: VarcharType => + new ArrayBasedMapBuilder(keyType, valueType).from( + new GenericArrayData(keys.toArray), new GenericArrayData(values.toArray)) + case _ => + // Preserve the historical behavior for ordinary string keys. + ArrayBasedMapData(keys.toArray, values.toArray) + } - if (badRecordException.isEmpty) { + // Ordinary value or key conversion failures invalidate the whole map. Delay throwing until + // the closing brace has been consumed and constrained-key deduplication has been applied. + badMapException.foreach(throw _) + + if (partialResultException.isEmpty) { mapData } else { - throw PartialMapDataResultException(mapData, badRecordException.get) + throw PartialMapDataResultException(mapData, partialResultException.get) } } @@ -717,6 +754,7 @@ class JacksonParser( } } catch { case e: SparkUpgradeException => throw e + case DuplicateMapKeyException(e) => throw e case e @ (_: RuntimeException | _: JsonProcessingException | _: MalformedInputException) => // JSON parser currently doesn't support partial results for corrupted records. // For such records, all fields other than the field configured by diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala index 00adcfe69bc56..c11dd3e02321c 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala @@ -26,6 +26,7 @@ import org.apache.spark.sql.catalyst.expressions.objects.StaticInvoke import org.apache.spark.sql.catalyst.parser.CatalystSqlParser import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String import org.apache.spark.util.ArrayImplicits._ object CharVarcharUtils extends Logging with SparkCharVarcharUtils { @@ -156,6 +157,25 @@ object CharVarcharUtils extends Logging with SparkCharVarcharUtils { StructType(fields) } + /** + * Applies CHAR padding and VARCHAR length checks when parsing text into a typed schema. + * Null stays null. Unbounded STRING is unchanged. This is assignment semantics + * (overflow raises EXCEED_LIMIT_LENGTH), not explicit CAST truncation. + */ + def applyTextParseSemantics(value: UTF8String, dt: DataType): UTF8String = { + if (value == null) { + null + } else { + dt match { + case c: CharType => + CharVarcharCodegenUtils.charTypeWriteSideCheck(value, c.length) + case v: VarcharType => + CharVarcharCodegenUtils.varcharTypeWriteSideCheck(value, v.length) + case _ => value + } + } + } + /** * Returns an expression to apply write-side string length check for the given expression. A * string value can not exceed N characters if it's written into a CHAR(N)/VARCHAR(N) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/FailureSafeParser.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/FailureSafeParser.scala index d9946d1b12ec3..f8e80657868e8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/FailureSafeParser.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/FailureSafeParser.scala @@ -17,11 +17,13 @@ package org.apache.spark.sql.catalyst.util +import org.apache.spark.SparkRuntimeException import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.GenericInternalRow import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.types.StructType import org.apache.spark.unsafe.types.UTF8String +import org.apache.spark.util.SparkErrorUtils class FailureSafeParser[IN]( rawParser: IN => Iterable[InternalRow], @@ -34,6 +36,13 @@ class FailureSafeParser[IN]( private val resultRow = new GenericInternalRow(schema.length) private val nullResult = new GenericInternalRow(schema.length) + private def duplicateMapKeyCause(e: Throwable): Option[SparkRuntimeException] = { + SparkErrorUtils.getRootCause(e) match { + case cause: SparkRuntimeException if cause.getCondition == "DUPLICATED_MAP_KEY" => Some(cause) + case _ => None + } + } + // This function takes 2 parameters: an optional partial result, and the bad record. If the given // schema doesn't contain a field for corrupted record, we just return the partial result or a // row with all fields null. If the given schema contains a field for corrupted record, we will @@ -59,6 +68,9 @@ class FailureSafeParser[IN]( try { rawParser.apply(input).iterator.map(row => toResultRow(Some(row), () => null)) } catch { + // Duplicate map keys are governed by mapKeyDedupPolicy, not the parse mode. + case e: BadRecordException if duplicateMapKeyCause(e).isDefined => + throw duplicateMapKeyCause(e).get case e: BadRecordException => mode match { case PermissiveMode => val partialResults = e.partialResults() diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/xml/StaxXmlParser.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/xml/StaxXmlParser.scala index 340dc61fb5112..63481f36a2cb9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/xml/StaxXmlParser.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/xml/StaxXmlParser.scala @@ -37,11 +37,11 @@ import com.google.common.io.ByteStreams import org.apache.hadoop.hdfs.BlockMissingException import org.apache.hadoop.security.AccessControlException -import org.apache.spark.{SparkIllegalArgumentException, SparkUpgradeException} +import org.apache.spark.{SparkIllegalArgumentException, SparkRuntimeException, SparkUpgradeException} import org.apache.spark.internal.Logging import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{ExprUtils, GenericInternalRow, ToStringBase} -import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, BadRecordException, DateFormatter, DropMalformedMode, FailureSafeParser, GenericArrayData, MapData, ParseMode, PartialResultArrayException, PartialResultException, PermissiveMode, TimeFormatter, TimestampFormatter} +import org.apache.spark.sql.catalyst.util.{ArrayBasedMapBuilder, ArrayBasedMapData, BadRecordException, CharVarcharUtils, DateFormatter, DropMalformedMode, FailureSafeParser, GenericArrayData, MapData, ParseMode, PartialResultArrayException, PartialResultException, PermissiveMode, TimeFormatter, TimestampFormatter} import org.apache.spark.sql.catalyst.util.LegacyDateFormats.FAST_DATE_FORMAT import org.apache.spark.sql.catalyst.xml.StaxXmlParser.convertStream import org.apache.spark.sql.errors.QueryExecutionErrors @@ -259,10 +259,14 @@ class StaxXmlParser( throw BadRecordException(xmlLiteral, () => Array.empty, wrappedCharException) case PartialResultException(row, cause) => - throw BadRecordException( - record = xmlLiteral, - partialResults = () => Array(row), - cause) + SparkErrorUtils.getRootCause(cause) match { + case e: SparkRuntimeException if e.getCondition == "DUPLICATED_MAP_KEY" => throw e + case _ => + throw BadRecordException( + record = xmlLiteral, + partialResults = () => Array(row), + cause) + } case PartialResultArrayException(rows, cause) => throw BadRecordException(record = xmlLiteral, partialResults = () => rows, cause) case e: Throwable => @@ -310,27 +314,27 @@ class StaxXmlParser( startElementName: String, attributes: Array[Attribute]): Any = dt match { case st: StructType => convertObject(parser, st) - case MapType(StringType, vt, _) => convertMap(parser, vt, attributes) + case MapType(kt: StringType, vt, _) => convertMap(parser, kt, vt, attributes) case ArrayType(st, _) => convertField(parser, st, startElementName) case VariantType => StaxXmlParser.convertVariant(parser, attributes, options) - case _: StringType => + case dt: StringType => convertTo( StaxXmlParserUtils.currentStructureAsString( parser, startElementName, options), - StringType) + dt) } (parser.peek, dataType) match { case (_: StartElement, dt: DataType) => convertComplicatedType(dt, startElementName, attributes) - case (_: EndElement, _: StringType) => + case (_: EndElement, dt: StringType) => StaxXmlParserUtils.skipNextEndElement(parser, startElementName, options) // Empty. It's null if "" is the null value if (options.nullValue == "") { null } else { - UTF8String.fromString("") + CharVarcharUtils.applyTextParseSemantics(UTF8String.fromString(""), dt) } case (_: EndElement, _: DataType) => StaxXmlParserUtils.skipNextEndElement(parser, startElementName, options) @@ -345,11 +349,11 @@ class StaxXmlParser( convertObject(parser, st) case (_: Characters, VariantType) => StaxXmlParser.convertVariant(parser, Array.empty, options) - case (_: Characters, _: StringType) => + case (_: Characters, dt: StringType) => convertTo( StaxXmlParserUtils.currentStructureAsString( parser, startElementName, options), - StringType) + dt) case (c: Characters, _: DataType) if c.isWhiteSpace => // When `Characters` is found, we need to look further to decide // if this is really data or space between other elements. @@ -374,31 +378,52 @@ class StaxXmlParser( */ private def convertMap( parser: XMLEventReader, + keyType: DataType, valueType: DataType, attributes: Array[Attribute]): MapData = { val kvPairs = ArrayBuffer.empty[(UTF8String, Any)] + var mapKeyException: Option[Throwable] = None + def mapKey(raw: String): UTF8String = { + CharVarcharUtils.applyTextParseSemantics(UTF8String.fromString(raw), keyType) + } + def appendPair(rawKey: String, value: Any): Unit = { + try { + kvPairs += (mapKey(rawKey) -> value) + } catch { + case NonFatal(e) => mapKeyException = mapKeyException.orElse(Some(e)) + } + } attributes.foreach { attr => - kvPairs += (UTF8String.fromString(options.attributePrefix + attr.getName.getLocalPart) - -> convertTo(attr.getValue, valueType)) + val value = convertTo(attr.getValue, valueType) + appendPair(options.attributePrefix + attr.getName.getLocalPart, value) } var shouldStop = false while (!shouldStop) { parser.nextEvent match { case e: StartElement => - val key = StaxXmlParserUtils.getName(e.asStartElement.getName, options) - kvPairs += - (UTF8String.fromString(key) -> convertField(parser, valueType, key)) + val rawKey = StaxXmlParserUtils.getName(e.asStartElement.getName, options) + val value = convertField(parser, valueType, rawKey) + appendPair(rawKey, value) case c: Characters if !c.isWhiteSpace => // Create a value tag field for it - kvPairs += // TODO: We don't support an array value tags in map yet. - (UTF8String.fromString(options.valueTag) -> convertTo(c.getData, valueType)) + val value = convertTo(c.getData, valueType) + appendPair(options.valueTag, value) case _: EndElement | _: EndDocument => shouldStop = true case _ => // do nothing } } - ArrayBasedMapData(kvPairs.toMap) + mapKeyException.foreach(throw _) + keyType match { + case _: CharType | _: VarcharType => + val mapBuilder = new ArrayBasedMapBuilder(keyType, valueType) + kvPairs.foreach { case (key, value) => mapBuilder.put(key, value) } + mapBuilder.build() + case _ => + // Preserve the historical last-wins behavior for ordinary string keys. + ArrayBasedMapData(kvPairs.toMap) + } } /** @@ -519,12 +544,12 @@ class StaxXmlParser( if (hasWildcard) { // Special case: there's an 'any' wildcard element that matches anything else // as a string (or array of strings, to parse multiple ones) - val newValue = convertField(parser, StringType, field) val anyIndex = schema.fieldIndex(wildcardColName) schema(wildcardColName).dataType match { - case StringType => - row(anyIndex) = newValue - case ArrayType(StringType, _) => + case dt: StringType => + row(anyIndex) = convertField(parser, dt, field) + case ArrayType(et: StringType, _) => + val newValue = convertField(parser, et, field) val values = Option(row(anyIndex)) .map(_.asInstanceOf[ArrayBuffer[String]]) .getOrElse(ArrayBuffer.empty[String]) @@ -536,6 +561,7 @@ class StaxXmlParser( } } catch { case e: SparkUpgradeException => throw e + case e: SparkRuntimeException if e.getCondition == "DUPLICATED_MAP_KEY" => throw e case NonFatal(e) => // TODO: we don't support partial results now badRecordException = badRecordException.orElse(Some(e)) @@ -606,7 +632,8 @@ class StaxXmlParser( timestampNTZFormatter.parseWithoutTimeZoneNanos(datum, t.precision, false) case _: DateType => parseXmlDate(datum, options) case _: TimeType => timeFormatter.parse(datum) - case _: StringType => UTF8String.fromString(datum) + case dt: StringType => + CharVarcharUtils.applyTextParseSemantics(UTF8String.fromString(datum), dt) case _: BinaryType => binaryParser(UTF8String.fromString(datum)) case _ => throw new SparkIllegalArgumentException( errorClass = "_LEGACY_ERROR_TEMP_3244", @@ -650,7 +677,7 @@ class StaxXmlParser( case LongType => signSafeToLong(value) case DoubleType => signSafeToDouble(value) case BooleanType => castTo(value, BooleanType) - case StringType => castTo(value, StringType) + case dt: StringType => castTo(value, dt) case BinaryType => castTo(value, BinaryType) case DateType => castTo(value, DateType) case TimestampType => castTo(value, TimestampType) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala index 39694ad3d6869..4d4e07f646e22 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CharVarcharTestSuite.scala @@ -845,6 +845,31 @@ trait CharVarcharTestSuite extends QueryTest { class BasicCharVarcharTestSuite extends SharedSparkSession { import testImplicits._ + private def assertParseExceedLimit(query: String): Unit = { + val e = intercept[SparkException] { sql(query).collect() } + val cause = e.getCause match { + case r: SparkRuntimeException => r + case other => + Option(other).flatMap(t => Option(t.getCause)).getOrElse(other) match { + case r: SparkRuntimeException => r + case _ => fail(s"expected EXCEED_LIMIT_LENGTH cause, got: $e") + } + } + checkError( + exception = cause, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "5")) + } + + private def assertDuplicateMapKey(query: String): Unit = { + checkError( + exception = intercept[SparkRuntimeException] { sql(query).collect() }, + condition = "DUPLICATED_MAP_KEY", + parameters = Map( + "key" -> "a ", + "mapKeyDedupPolicy" -> "\"spark.sql.mapKeyDedupPolicy\"")) + } + test("user-specified schema in cast") { def assertNoCharType(df: DataFrame): Unit = { checkAnswer(df, Row("0")) @@ -2226,6 +2251,138 @@ class BasicCharVarcharTestSuite extends SharedSparkSession { } } } + + test("SPARK-59274: from_json/csv/xml honor CHAR/VARCHAR under standardSemantics") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val jsonChar = sql("""SELECT from_json('{"a": "str"}', 'a CHAR(5)')""") + val jsonCharType = jsonChar.schema.head.dataType.asInstanceOf[StructType] + assert(jsonCharType.head.dataType === CharType(5)) + checkAnswer(jsonChar, Row(Row("str "))) + + val jsonVarchar = sql("""SELECT from_json('{"a": "ab"}', 'a VARCHAR(5)')""") + val jsonVarcharType = jsonVarchar.schema.head.dataType.asInstanceOf[StructType] + assert(jsonVarcharType.head.dataType === VarcharType(5)) + checkAnswer(jsonVarchar, Row(Row("ab"))) + + // Default PERMISSIVE mode turns length failures into a null record. + Seq("CHAR(5)", "VARCHAR(5)").foreach { dataType => + checkAnswer( + sql(s"""SELECT from_json('{"a": "abcdef"}', 'a $dataType')"""), + Row(Row(null))) + assertParseExceedLimit( + s"""SELECT from_json( + | '{"a": "abcdef"}', + | 'a $dataType', + | map('mode', 'FAILFAST'))""".stripMargin) + } + + checkAnswer( + sql("""SELECT from_json('{"ab": 1}', 'MAP')"""), + Row(Map("ab " -> 1))) + + checkAnswer(sql("SELECT from_csv('str', 'a CHAR(5)')"), Row(Row("str "))) + Seq("CHAR(5)", "VARCHAR(5)").foreach { dataType => + checkAnswer(sql(s"SELECT from_csv('abcdef', 'a $dataType')"), Row(Row(null))) + assertParseExceedLimit( + s"SELECT from_csv('abcdef', 'a $dataType', map('mode', 'FAILFAST'))") + } + + checkAnswer( + sql("SELECT from_xml('str', 'a CHAR(5)')"), + Row(Row("str "))) + checkAnswer( + sql( + """SELECT from_xml( + | '', + | 'a CHAR(5)', + | map('nullValue', 'NULL'))""".stripMargin), + Row(Row(" "))) + Seq("CHAR(5)", "VARCHAR(5)").foreach { dataType => + checkAnswer( + sql(s"SELECT from_xml('abcdef', 'a $dataType')"), + Row(Row(null))) + assertParseExceedLimit( + s"SELECT from_xml('abcdef', 'a $dataType', " + + "map('mode', 'FAILFAST'))") + } + checkAnswer( + sql("SELECT from_xml('1', 'm MAP')"), + Row(Row(Map("ab " -> 1)))) + + checkAnswer( + sql("""SELECT schema_of_json(CAST('{"a":1}' AS VARCHAR(20)))"""), + Row("STRUCT")) + checkAnswer( + sql("SELECT schema_of_csv(CAST('1,abc' AS VARCHAR(20)))"), + Row("STRUCT<_c0: INT, _c1: STRING>")) + checkAnswer( + sql("SELECT schema_of_xml(CAST('1' AS VARCHAR(40)))"), + Row("STRUCT")) + } + } + + test("SPARK-59274: normalized CHAR map key collisions honor the dedup policy") { + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true", + SQLConf.JSON_ENABLE_PARTIAL_RESULTS.key -> "true") { + val jsonQuery = + """SELECT from_json('{"a":1,"a ":2}', 'MAP')""" + val nestedJsonQuery = + """SELECT from_json( + | '{"outer":{"a":1,"a ":2}}', + | 'MAP>')""".stripMargin + val badFieldBeforeDuplicateQuery = + """SELECT from_json( + | '{"bad":"not-an-int","m":{"a":1,"a ":2}}', + | 'bad INT, m MAP')""".stripMargin + val badKeyThenSiblingQuery = + """SELECT from_json( + | '{"m":{"abc":1},"tail":2}', + | 'm MAP, tail INT').tail""".stripMargin + val badXmlKeyThenSiblingQuery = + """SELECT from_xml( + | '12', + | 'm MAP, tail INT').tail""".stripMargin + val badValueBeforeDuplicateQuery = + """SELECT from_json('{"bad":"not-an-int","a":1,"a ":2}', 'MAP')""" + val xmlQuery = + """SELECT from_xml( + | '19', + | 'm MAP', + | map('valueTag', 'a ')).m""".stripMargin + + assertDuplicateMapKey(jsonQuery) + assertDuplicateMapKey(nestedJsonQuery) + assertDuplicateMapKey(badFieldBeforeDuplicateQuery) + assertDuplicateMapKey(badValueBeforeDuplicateQuery) + assertDuplicateMapKey(xmlQuery) + checkAnswer(sql(badKeyThenSiblingQuery), Row(2)) + checkAnswer(sql(badXmlKeyThenSiblingQuery), Row(2)) + + withSQLConf( + SQLConf.MAP_KEY_DEDUP_POLICY.key -> SQLConf.MapKeyDedupPolicy.LAST_WIN.toString) { + checkAnswer(sql(jsonQuery), Row(Map("a " -> 2))) + checkAnswer(sql(nestedJsonQuery), Row(Map("outer" -> Map("a " -> 2)))) + checkAnswer(sql(badValueBeforeDuplicateQuery), Row(null)) + checkAnswer(sql(xmlQuery), Row(Map("a " -> 9))) + } + } + } + + test("SPARK-59274: ordinary STRING map duplicate behavior is unchanged") { + Seq("false", "true").foreach { standardSemantics => + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> standardSemantics) { + checkAnswer( + sql("""SELECT from_json('{"a":1,"a":2}', 'MAP')"""), + Row(Map("a" -> 2))) + checkAnswer( + sql("""SELECT from_xml( + | '12', + | 'm MAP').m""".stripMargin), + Row(Map("a" -> 2))) + } + } + } } class FileSourceCharVarcharTestSuite extends CharVarcharTestSuite with SharedSparkSession {