From f4d6fbbc5c9ec3c6607ee6334607384a8e9b50f0 Mon Sep 17 00:00:00 2001 From: Pedrum Jalali Date: Wed, 16 Sep 2026 10:13:25 -0700 Subject: [PATCH] [MINOR][SQL][TEST] Make the not-null and null-map-key test assertions overridable ### What changes were proposed in this pull request? Refactors test assertions into protected methods. ### Why are the changes needed? Apache Gluten reuses `RuntimeNullChecksV2Writes` and `DataFrameFunctionsSuite` against its Velox backend. As an example in apache/gluten#12976, after the map_from_arrays operator was enabled, it changed the exception thrown in these tests and required copying over the entire test only to override the exception type in the assertion. By making the assertion protected downstream consumers can override the limited lines without having to copy the entire test body allowing better re-usability downstream. ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? ``` build/mvn -pl sql/core -am test -Dtest=none -DfailIfNoTests=false \ -DwildcardSuites=org.apache.spark.sql.RuntimeNullChecksV2Writes,org.apache.spark.sql.DataFrameFunctionsSuite RuntimeNullChecksV2Writes: ... DataFrameFunctionsSuite: ... Run completed in 50 seconds, 854 milliseconds. Total number of tests run: 176 Suites: completed 4, aborted 0 Tests: succeeded 176, failed 0, canceled 0, ignored 0, pending 0 All tests passed. [INFO] BUILD SUCCESS ``` ### Was this patch authored or co-authored using generative AI tooling? Generated-by: Co-authored with Claude Code (Opus 5) --- .../spark/sql/DataFrameFunctionsSuite.scala | 48 ++++++++----------- .../spark/sql/RuntimeNullChecksV2Writes.scala | 35 +++++--------- 2 files changed, 33 insertions(+), 50 deletions(-) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameFunctionsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameFunctionsSuite.scala index c29b993e500c1..f905ecf2aba0e 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameFunctionsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameFunctionsSuite.scala @@ -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 @@ -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] { @@ -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] { @@ -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") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/RuntimeNullChecksV2Writes.scala b/sql/core/src/test/scala/org/apache/spark/sql/RuntimeNullChecksV2Writes.scala index f384f892299e8..03c6076384c56 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/RuntimeNullChecksV2Writes.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/RuntimeNullChecksV2Writes.scala @@ -53,7 +53,7 @@ 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() @@ -61,7 +61,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { sql("INSERT INTO t VALUES ('txt', null)") } } - assert(e.getCondition == "NOT_NULL_ASSERT_VIOLATION") } } @@ -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 @@ -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 @@ -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 @@ -137,7 +134,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { """.stripMargin) } } - assertNotNullException(e3, Seq("s", "m")) } } @@ -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 @@ -190,7 +186,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { """.stripMargin) } } - assertNotNullException(e, Seq("s", "ns", "x")) } } @@ -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 @@ -239,7 +234,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { """.stripMargin) } } - assertNotNullException(e, Seq("arr", "element")) } } @@ -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 @@ -295,7 +289,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { """.stripMargin) } } - assertNotNullException(e, Seq("arr", "element", "x")) } } @@ -326,7 +319,7 @@ 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() @@ -334,7 +327,6 @@ class RuntimeNullChecksV2Writes extends SharedSparkSession { sql("INSERT INTO t VALUES (1 AS i, map(1, null) AS m)") } } - assertNotNullException(e, Seq("m", "value")) } } @@ -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 @@ -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 @@ -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") } } }