Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ import org.apache.spark.sql.catalyst.expressions.{
Alias,
Attribute,
BinaryComparison,
Cast,
Expression,
In,
Literal,
Expand Down Expand Up @@ -48,6 +49,19 @@ object ApplyCharTypePaddingHelper {
}
}

private object StringAttrOrOuterRef {
def unapply(e: Expression): Option[(Expression, Attribute)] = e match {
case attr @ AttrOrOuterRef(a) if attr.dataType.isInstanceOf[StringType] =>
Some(attr -> a)
case cast @ Cast(AttrOrOuterRef(a), _: StringType, _, _)
if a.dataType.isInstanceOf[StringType] &&
!cast.containsTag(Cast.USER_SPECIFIED_CAST) =>
Some(cast -> a)
case _ =>
None
}
}

private[sql] def readSidePadding(
relation: LogicalPlan,
cleanedRelation: () => LogicalPlan): (LogicalPlan, Seq[(Attribute, Attribute)]) = {
Expand Down Expand Up @@ -98,23 +112,25 @@ object ApplyCharTypePaddingHelper {
}
.getOrElse(b)

case i @ In(e @ AttrOrOuterRef(attr), list)
if i.resolved && attr.dataType.isInstanceOf[StringType] && list.forall(_.foldable) =>
case i @ In(StringAttrOrOuterRef(e, attr), list)
if i.resolved && list.forall(_.foldable) =>
CharVarcharUtils
.getRawType(attr.metadata)
.flatMap {
case c: CharType =>
val (nulls, literalChars) =
list.map(_.eval().asInstanceOf[UTF8String]).partition(_ == null)
val literalCharLengths = literalChars.map(_.numChars())
val literalValues = list.map(lit => lit -> lit.eval().asInstanceOf[UTF8String])
val literalCharLengths = literalValues.collect {
case (_, value) if value != null => value.numChars()
}
val targetLen = (c.length +: literalCharLengths).max
Some(
i.copy(
value = addPadding(e, c.length, targetLen, alwaysPad = padCharCol),
list = list.zip(literalCharLengths).map {
case (lit, charLength) =>
addPadding(lit, charLength, targetLen, alwaysPad = false)
} ++ nulls.map(Literal.create(_, StringType))
list = literalValues.map {
case (lit, null) => lit
case (lit, value) =>
addPadding(lit, value.numChars(), targetLen, alwaysPad = false)
}
)
)
case _ => None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -337,12 +337,21 @@ object CharVarcharUtils extends Logging with SparkCharVarcharUtils {
val fieldExpr = GetStructField(expr, i, Some(field.name))
val padded = padCharToTargetLength(
fieldExpr, field.dataType, targets(i).dataType, alwaysPad)
needPadding = padded.isDefined
needPadding = needPadding || padded.isDefined
createStructExprs += Literal(field.name)
createStructExprs += padded.getOrElse(fieldExpr)
i += 1
}
if (needPadding) Some(CreateNamedStruct(createStructExprs.toSeq)) else None
if (needPadding) {
val struct = CreateNamedStruct(createStructExprs.toSeq)
if (expr.nullable) {
Some(If(IsNull(expr), Literal(null, struct.dataType), struct))
} else {
Some(struct)
}
} else {
None
}

case (ArrayType(et, containsNull), ArrayType(target, _)) =>
val param = NamedLambdaVariable("x", replaceCharVarcharWithString(et), containsNull)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -681,6 +681,52 @@ trait CharVarcharTestSuite extends QueryTest {
}
}

test("SPARK-59278: CHAR padding preserves IN list order around NULL") {
withTable("t") {
sql(s"CREATE TABLE t(c CHAR(5)) USING $format")
sql("INSERT INTO t VALUES ('a')")

checkAnswer(sql("SELECT c IN (NULL, 'a') FROM t"), Row(true))
checkAnswer(sql("SELECT c IN ('x', NULL, 'a') FROM t"), Row(true))
checkAnswer(sql("SELECT c IN (NULL, 'x') FROM t"), Row(null))
checkAnswer(sql("SELECT c NOT IN (NULL, 'a') FROM t"), Row(false))
checkAnswer(sql("SELECT c NOT IN (NULL, 'x') FROM t"), Row(null))

withSQLConf(SQLConf.OPTIMIZER_INSET_CONVERSION_THRESHOLD.key -> "1") {
checkAnswer(sql("SELECT c IN (NULL, 'a', 'b') FROM t"), Row(true))
checkAnswer(sql("SELECT c IN (NULL, 'x', 'y') FROM t"), Row(null))
}

// A user-specified cast deliberately opts into STRING comparison and must not trigger the
// implicit CHAR padding rewrite.
checkAnswer(sql("SELECT CAST(c AS STRING) IN ('a') FROM t"), Row(false))
checkAnswer(sql("SELECT CAST(c AS STRING) IN (NULL, 'a') FROM t"), Row(null))
}
}

test("SPARK-59278: CHAR padding accumulates across nested struct fields") {
withTable("t1", "t2") {
sql(s"CREATE TABLE t1(s STRUCT<a: CHAR(2), b: CHAR(5)>) USING $format")
sql(s"CREATE TABLE t2(s STRUCT<a: CHAR(4), b: CHAR(5)>) USING $format")
sql("INSERT INTO t1 SELECT named_struct('a', 'a', 'b', 'b')")
sql("INSERT INTO t2 SELECT named_struct('a', 'a', 'b', 'b')")

checkAnswer(sql("SELECT t1.s = t2.s FROM t1 CROSS JOIN t2"), Row(true))

sql("INSERT OVERWRITE t1 SELECT CAST(NULL AS STRUCT<a: STRING, b: STRING>)")
sql("INSERT OVERWRITE t2 SELECT CAST(NULL AS STRUCT<a: STRING, b: STRING>)")
checkAnswer(sql("SELECT t1.s = t2.s FROM t1 CROSS JOIN t2"), Row(null))
checkAnswer(sql("SELECT t1.s <=> t2.s FROM t1 CROSS JOIN t2"), Row(true))

sql(
"""INSERT OVERWRITE t2
|SELECT named_struct(
| 'a', CAST(NULL AS STRING),
| 'b', CAST(NULL AS STRING))""".stripMargin)
checkAnswer(sql("SELECT t1.s <=> t2.s FROM t1 CROSS JOIN t2"), Row(false))
}
}

test("SPARK-35359: create table and insert data over length values") {
Seq("char", "varchar").foreach { typ =>
withSQLConf((SQLConf.LEGACY_CHAR_VARCHAR_AS_STRING.key, "true")) {
Expand Down