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 @@ -45,6 +45,14 @@ import org.apache.spark.tags.ExtendedSQLTest
class DataFrameFunctionsSuite extends SharedSparkSession {
import testImplicits._

protected def assertNullMapKeyFailure(func: => Any): Unit = {
checkError(
exception = intercept[SparkRuntimeException](func),
condition = "NULL_MAP_KEY",
parameters = Map.empty
)
}

test("DataFrame function and SQL function parity") {
// This test compares the available list of DataFrame functions in
// org.apache.spark.sql.functions with the SQL function registry. This attempts to verify that
Expand Down Expand Up @@ -186,13 +194,9 @@ class DataFrameFunctionsSuite extends SharedSparkSession {
)

val df5 = Seq((Seq("a", null), Seq(1, 2))).toDF("k", "v")
checkError(
exception = intercept[SparkRuntimeException] {
df5.select(map_from_arrays($"k", $"v")).collect()
},
condition = "NULL_MAP_KEY",
parameters = Map.empty
)
assertNullMapKeyFailure {
df5.select(map_from_arrays($"k", $"v")).collect()
}

val df6 = Seq((Seq(1, 2), Seq("a"))).toDF("k", "v")
val msg2 = intercept[Exception] {
Expand Down Expand Up @@ -5432,21 +5436,13 @@ class DataFrameFunctionsSuite extends SharedSparkSession {
stop = 35)
)

checkError(
exception = intercept[SparkRuntimeException] {
dfExample1.selectExpr("transform_keys(i, (k, v) -> v)").show()
},
condition = "NULL_MAP_KEY",
parameters = Map.empty
)
assertNullMapKeyFailure {
dfExample1.selectExpr("transform_keys(i, (k, v) -> v)").show()
}

checkError(
exception = intercept[SparkRuntimeException] {
dfExample1.select(transform_keys(col("i"), (k, v) => v)).show()
},
condition = "NULL_MAP_KEY",
parameters = Map.empty
)
assertNullMapKeyFailure {
dfExample1.select(transform_keys(col("i"), (k, v) => v)).show()
}

checkError(
exception = intercept[AnalysisException] {
Expand Down Expand Up @@ -6070,13 +6066,9 @@ class DataFrameFunctionsSuite extends SharedSparkSession {

test("SPARK-24734: Fix containsNull of Concat for array type") {
val df = Seq((Seq(1), Seq[Integer](null), Seq("a", "b"))).toDF("k1", "k2", "v")
checkError(
exception = intercept[SparkRuntimeException] {
df.select(map_from_arrays(concat($"k1", $"k2"), $"v")).show()
},
condition = "NULL_MAP_KEY",
parameters = Map.empty
)
assertNullMapKeyFailure {
df.select(map_from_arrays(concat($"k1", $"k2"), $"v")).show()
}
}

test("SPARK-26370: Fix resolution of higher-order function for the same identifier") {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -53,15 +53,14 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
withTable("t") {
sql(s"CREATE TABLE t (s STRING, i INT NOT NULL) USING $FORMAT")

val e = intercept[SparkRuntimeException] {
assertNotNullException(Seq("i")) {
if (byName) {
val inputDF = sql("SELECT 'txt' AS s, null AS i")
inputDF.writeTo("t").append()
} else {
sql("INSERT INTO t VALUES ('txt', null)")
}
}
assert(e.getCondition == "NOT_NULL_ASSERT_VIOLATION")
}
}

Expand All @@ -85,7 +84,7 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
|USING $FORMAT
""".stripMargin)

val e1 = intercept[SparkRuntimeException] {
assertNotNullException(Seq("s", "ns")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -101,9 +100,8 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e1, Seq("s", "ns"))

val e2 = intercept[SparkRuntimeException] {
assertNotNullException(Seq("s", "arr")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -119,9 +117,8 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e2, Seq("s", "arr"))

val e3 = intercept[SparkRuntimeException] {
assertNotNullException(Seq("s", "m")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -137,7 +134,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e3, Seq("s", "m"))
}
}

Expand Down Expand Up @@ -174,7 +170,7 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
}
checkAnswer(spark.table("t"), Row(1, Row(1, null)))

val e = intercept[SparkRuntimeException] {
assertNotNullException(Seq("s", "ns", "x")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -190,7 +186,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e, Seq("s", "ns", "x"))
}
}

Expand Down Expand Up @@ -223,7 +218,7 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
}
checkAnswer(spark.table("t"), Row(1, null))

val e = intercept[SparkRuntimeException] {
assertNotNullException(Seq("arr", "element")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -239,7 +234,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e, Seq("arr", "element"))
}
}

Expand Down Expand Up @@ -279,7 +273,7 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
}
checkAnswer(spark.table("t"), Row(1, List(null, Row(1, 1))))

val e = intercept[SparkRuntimeException] {
assertNotNullException(Seq("arr", "element", "x")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -295,7 +289,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e, Seq("arr", "element", "x"))
}
}

Expand Down Expand Up @@ -326,15 +319,14 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
}
checkAnswer(spark.table("t"), Row(1, null))

val e = intercept[SparkRuntimeException] {
assertNotNullException(Seq("m", "value")) {
if (byName) {
val inputDF = sql("SELECT 1 AS i, map(1, null) AS m")
inputDF.writeTo("t").append()
} else {
sql("INSERT INTO t VALUES (1 AS i, map(1, null) AS m)")
}
}
assertNotNullException(e, Seq("m", "value"))
}
}

Expand Down Expand Up @@ -366,7 +358,7 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
}
checkAnswer(spark.table("t"), Row(1, Map(Row(1, 1) -> null)))

val e1 = intercept[SparkRuntimeException] {
assertNotNullException(Seq("m", "key", "x")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -382,9 +374,8 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e1, Seq("m", "key", "x"))

val e2 = intercept[SparkRuntimeException] {
assertNotNullException(Seq("m", "value", "x")) {
if (byName) {
val inputDF = sql(
s"""SELECT
Expand All @@ -400,15 +391,15 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession {
""".stripMargin)
}
}
assertNotNullException(e2, Seq("m", "value", "x"))
}
}

private def assertNotNullException(e: SparkRuntimeException, colPath: Seq[String]): Unit = {
protected def assertNotNullException(colPath: Seq[String])(func: => Any): Unit = {
val e = intercept[SparkRuntimeException](func)
e.getCause match {
case _ if e.getCondition == "NOT_NULL_ASSERT_VIOLATION" =>
case other =>
fail(s"Unexpected exception cause: $other")
fail(s"Unexpected exception cause for ${colPath.mkString(".")}: $other")
}
}
}