From 8c7e7c78eee474f2bb5aedb45bdd487a7d132efc Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sat, 19 Sep 2026 04:24:41 +0000 Subject: [PATCH 1/3] fix: defer throwing literal casts to Spark at runtime --- .../apache/comet/expressions/CometCast.scala | 18 ++++++-- .../serde/CometScalarFunctionSuite.scala | 42 ++++++++++++++++++- 2 files changed, 54 insertions(+), 6 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala index e14fb086b12..5f279dbda49 100644 --- a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala +++ b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala @@ -19,6 +19,8 @@ package org.apache.comet.expressions +import scala.util.control.NonFatal + import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Expression, Literal} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, DataType, DataTypes, DecimalType, MapType, NullType, StructType, TimestampNTZType, TimestampType} @@ -98,9 +100,9 @@ object CometCast override def getSupportLevel(cast: Cast): SupportLevel = { if (cast.child.isInstanceOf[Literal]) { - // A cast whose child is a literal is folded by Spark at planning time via `cast.eval()` - // (see `convert`), so the cast never executes natively and the result matches Spark by - // definition. `CometLiteral` then validates the resulting literal's data type, except + // A successfully evaluated literal cast is folded via `cast.eval()` (see `convert`). + // Spark can leave a throwing literal cast in an unvisited conditional branch, so `convert` + // falls back when evaluation fails. `CometLiteral` validates a folded literal, except // for `VariantType` which must be rejected here: the fold produces a `Literal[VariantType]` // that no downstream Comet serde can serialize. if (isVariantType(cast.child.dataType) || isVariantType(cast.dataType)) { @@ -121,7 +123,15 @@ object CometCast val cometEvalMode = evalMode(cast) cast.child match { case _: Literal => - exprToProtoInternal(Literal.create(cast.eval(), cast.dataType), inputs, binding) + val value = + try { + cast.eval() + } catch { + case NonFatal(_) => + withFallbackReason(cast, "Literal cast requires Spark's conditional evaluation") + return None + } + exprToProtoInternal(Literal.create(value, cast.dataType), inputs, binding) case _ => if (isAlwaysCastToNull(cast.child.dataType, cast.dataType, cometEvalMode)) { exprToProtoInternal(Literal.create(null, cast.dataType), inputs, binding) diff --git a/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala b/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala index 8cf88ec5593..220176aac0f 100644 --- a/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala @@ -19,11 +19,14 @@ package org.apache.comet.serde -import org.apache.spark.sql.CometTestBase +import org.apache.spark.SparkThrowable +import org.apache.spark.sql.{CometTestBase, Row} import org.apache.spark.sql.catalyst.expressions.{Abs, Cos, Expression, Literal, Round, Unevaluable} +import org.apache.spark.sql.comet.CometNativeScanExec +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, DoubleType, IntegerType} -import org.apache.comet.{CometExplainInfo, CometSparkSessionExtensions} +import org.apache.comet.{CometConf, CometExplainInfo, CometSparkSessionExtensions} /** * Synthetic expression whose constructor declares `evalMode`, used to prove class-level detection @@ -216,6 +219,41 @@ class CometScalarFunctionSuite extends CometTestBase { assertRejectReason(withContext, "CometScalarFunction", "evalContext") } + test("literal cast failure in an unvisited conditional branch does not fail planning") { + withTempPath { path => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(2) + .selectExpr("id", "CAST(id AS STRING) AS value") + .coalesce(1) + .write + .parquet(path.getCanonicalPath) + } + withSQLConf(SQLConf.ANSI_ENABLED.key -> "true") { + withParquetTable(path.getCanonicalPath, "cast_branch_rows") { + val cast = "CAST(IF(id = 1, 'bad', value) AS INT)" + val masked = s"SELECT $cast AS parsed FROM cast_branch_rows LIMIT 1" + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + assert(sql(masked).collect().toSeq == Seq(Row(0))) + } + val (_, plan) = checkSparkAnswerAndFallbackReason( + masked, + "Literal cast requires Spark's conditional evaluation") + assert(collect(plan) { case scan: CometNativeScanExec => scan }.nonEmpty) + val (sparkError, cometError) = + checkSparkAnswerMaybeThrows(sql(s"SELECT $cast FROM cast_branch_rows")) + assert(sparkError.nonEmpty && cometError.nonEmpty) + val errors = Seq(sparkError.get, cometError.get).map { error => + causeChain(error).collect { case e: SparkThrowable => e }.last + } + assert(errors.forall(_.getErrorClass == "CAST_INVALID_INPUT")) + assert(errors(0).getClass == errors(1).getClass) + assert(errors(0).getSqlState == errors(1).getSqlState) + } + } + } + } + test("CometScalarFunction allows non-ANSI expressions") { val cos = Cos(Literal(0.0)) val result = CometScalarFunction[Cos]("cos").convert(cos, Seq.empty, binding = true) From 73eb76d4f6506d4e4e5afc0468e34eaf26975405 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sun, 4 Oct 2026 17:33:46 +0000 Subject: [PATCH 2/3] test: cover deliberate literal cast fallback and native controls --- .../apache/comet/expressions/CometCast.scala | 10 ++- .../apache/comet/CometNativeCastSuite.scala | 64 ++++++++++++++++++- .../serde/CometScalarFunctionSuite.scala | 42 +----------- 3 files changed, 74 insertions(+), 42 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala index 5f279dbda49..a3193931f06 100644 --- a/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala +++ b/spark/src/main/scala/org/apache/comet/expressions/CometCast.scala @@ -43,6 +43,10 @@ object CometCast private[comet] val negativeScaleDecimalToStringReason: String = "Negative-scale decimal requires spark.sql.legacy.allowNegativeScaleOfDecimal=true" + private[comet] val literalCastConditionalEvalReason: String = + "Cast of a literal threw during planning; Spark leaves it for conditional evaluation " + + "so it may never be reached at runtime." + // When `spark.sql.legacy.castComplexTypesToString.enabled` is true, Spark wraps maps and // structs with `[]` (instead of `{}`) when casting to string, and omits NULL elements of // structs/maps/arrays (instead of rendering them as the literal "null"). Comet's native @@ -128,7 +132,11 @@ object CometCast cast.eval() } catch { case NonFatal(_) => - withFallbackReason(cast, "Literal cast requires Spark's conditional evaluation") + // ConstantFolding.tryFold leaves failed conditional expressions unfolded. Its + // FAILED_TO_EVALUATE tag is private[sql], so evaluate defensively here instead. + // Deliberately keep this in convert: a codegen-dispatched projection evaluates + // a whole batch and can reach a throwing row that Spark skips under LIMIT. + withFallbackReason(cast, literalCastConditionalEvalReason) return None } exprToProtoInternal(Literal.create(value, cast.dataType), inputs, binding) diff --git a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala index 4db950e73c9..8cce831b981 100644 --- a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala @@ -29,8 +29,9 @@ import scala.util.Random import org.apache.hadoop.fs.Path import org.apache.spark.SparkConf import org.apache.spark.sql.{CometTestBase, DataFrame, Row, SaveMode} -import org.apache.spark.sql.catalyst.expressions.Cast +import org.apache.spark.sql.catalyst.expressions.{Cast, Literal} import org.apache.spark.sql.catalyst.parser.ParseException +import org.apache.spark.sql.comet.{CometNativeScanExec, CometProjectExec} import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.functions.{col, monotonically_increasing_id} import org.apache.spark.sql.internal.SQLConf @@ -65,6 +66,67 @@ class CometNativeCastSuite extends CometTestBase with AdaptiveSparkPlanHelper { import testImplicits._ + test("literal cast failure in an unvisited conditional branch does not fail planning") { + withTempPath { path => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(2) + .selectExpr("id", "CAST(id AS STRING) AS value") + .coalesce(1) + .write + .parquet(path.getCanonicalPath) + } + withSQLConf(SQLConf.ANSI_ENABLED.key -> "true") { + withParquetTable(path.getCanonicalPath, "cast_branch_rows") { + val cast = "CAST(IF(id = 1, 'bad', value) AS INT)" + val masked = s"SELECT $cast AS parsed FROM cast_branch_rows LIMIT 1" + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + assert(sql(masked).collect().toSeq == Seq(Row(0))) + } + // ConstantFolding must leave the throwing literal cast inside the conditional. + val badCasts = sql(masked).queryExecution.optimizedPlan + .flatMap(_.expressions) + .flatMap(_.collect { + case c: Cast + if c.child.isInstanceOf[Literal] && + Option(c.child.asInstanceOf[Literal].value).exists(_.toString == "bad") => + c + }) + assert(badCasts.nonEmpty) + val (_, plan) = + checkSparkAnswerAndFallbackReason(masked, CometCast.literalCastConditionalEvalReason) + assert(collect(plan) { case scan: CometNativeScanExec => scan }.nonEmpty) + checkSparkError(sql(s"SELECT $cast FROM cast_branch_rows"), "CAST_INVALID_INPUT") + } + } + } + } + + test("successful literal casts and non-ANSI invalid casts retain native projection") { + withParquetTable(Seq(0, 1), "cast_branch_controls") { + for (ansi <- Seq("true", "false")) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { + val query = "SELECT CAST(IF(value = 1, '1', CAST(value AS STRING)) AS INT) " + + "FROM cast_branch_controls" + val (_, plan) = checkSparkAnswerAndOperator(sql(query)) + assert(collect(plan) { case project: CometProjectExec => project }.nonEmpty) + assert( + !plan.exists(_.getTagValue(CometExplainInfo.FALLBACK_REASONS) + .exists(_.contains(CometCast.literalCastConditionalEvalReason)))) + } + } + withSQLConf(SQLConf.ANSI_ENABLED.key -> "false") { + val query = "SELECT CAST(IF(value = 1, 'bad', CAST(value AS STRING)) AS INT) " + + "FROM cast_branch_controls" + val (_, plan) = checkSparkAnswerAndOperator(sql(query)) + assert(collect(plan) { case project: CometProjectExec => project }.nonEmpty) + assert( + !plan.exists(_.getTagValue(CometExplainInfo.FALLBACK_REASONS) + .exists(_.contains(CometCast.literalCastConditionalEvalReason)))) + } + } + } + // Casts in this suite predominantly test non-ANSI semantics (silent overflow/null on // invalid input); tests that target ANSI behavior opt in explicitly via withSQLConf. override protected def sparkConf: SparkConf = diff --git a/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala b/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala index 220176aac0f..8cf88ec5593 100644 --- a/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/serde/CometScalarFunctionSuite.scala @@ -19,14 +19,11 @@ package org.apache.comet.serde -import org.apache.spark.SparkThrowable -import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.catalyst.expressions.{Abs, Cos, Expression, Literal, Round, Unevaluable} -import org.apache.spark.sql.comet.CometNativeScanExec -import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, DoubleType, IntegerType} -import org.apache.comet.{CometConf, CometExplainInfo, CometSparkSessionExtensions} +import org.apache.comet.{CometExplainInfo, CometSparkSessionExtensions} /** * Synthetic expression whose constructor declares `evalMode`, used to prove class-level detection @@ -219,41 +216,6 @@ class CometScalarFunctionSuite extends CometTestBase { assertRejectReason(withContext, "CometScalarFunction", "evalContext") } - test("literal cast failure in an unvisited conditional branch does not fail planning") { - withTempPath { path => - withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - spark - .range(2) - .selectExpr("id", "CAST(id AS STRING) AS value") - .coalesce(1) - .write - .parquet(path.getCanonicalPath) - } - withSQLConf(SQLConf.ANSI_ENABLED.key -> "true") { - withParquetTable(path.getCanonicalPath, "cast_branch_rows") { - val cast = "CAST(IF(id = 1, 'bad', value) AS INT)" - val masked = s"SELECT $cast AS parsed FROM cast_branch_rows LIMIT 1" - withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - assert(sql(masked).collect().toSeq == Seq(Row(0))) - } - val (_, plan) = checkSparkAnswerAndFallbackReason( - masked, - "Literal cast requires Spark's conditional evaluation") - assert(collect(plan) { case scan: CometNativeScanExec => scan }.nonEmpty) - val (sparkError, cometError) = - checkSparkAnswerMaybeThrows(sql(s"SELECT $cast FROM cast_branch_rows")) - assert(sparkError.nonEmpty && cometError.nonEmpty) - val errors = Seq(sparkError.get, cometError.get).map { error => - causeChain(error).collect { case e: SparkThrowable => e }.last - } - assert(errors.forall(_.getErrorClass == "CAST_INVALID_INPUT")) - assert(errors(0).getClass == errors(1).getClass) - assert(errors(0).getSqlState == errors(1).getSqlState) - } - } - } - } - test("CometScalarFunction allows non-ANSI expressions") { val cos = Cos(Literal(0.0)) val result = CometScalarFunction[Cos]("cos").convert(cos, Seq.empty, binding = true) From cae771a9db078dac28473841ceb7e15c44e02990 Mon Sep 17 00:00:00 2001 From: Chao Sun Date: Sun, 4 Oct 2026 17:50:07 +0000 Subject: [PATCH 3/3] test: exercise cast fallback with shared Spark error checks --- .../scala/org/apache/comet/CometNativeCastSuite.scala | 11 +++++++---- .../scala/org/apache/spark/sql/CometTestBase.scala | 9 ++++++--- 2 files changed, 13 insertions(+), 7 deletions(-) diff --git a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala index 8cce831b981..73949a2c320 100644 --- a/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala @@ -96,17 +96,20 @@ class CometNativeCastSuite extends CometTestBase with AdaptiveSparkPlanHelper { val (_, plan) = checkSparkAnswerAndFallbackReason(masked, CometCast.literalCastConditionalEvalReason) assert(collect(plan) { case scan: CometNativeScanExec => scan }.nonEmpty) - checkSparkError(sql(s"SELECT $cast FROM cast_branch_rows"), "CAST_INVALID_INPUT") + checkSparkError( + sql(s"SELECT $cast FROM cast_branch_rows"), + "CAST_INVALID_INPUT", + checkNative = false) } } } } test("successful literal casts and non-ANSI invalid casts retain native projection") { - withParquetTable(Seq(0, 1), "cast_branch_controls") { + withParquetTable(Seq(Tuple1(0), Tuple1(1)), "cast_branch_controls") { for (ansi <- Seq("true", "false")) { withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { - val query = "SELECT CAST(IF(value = 1, '1', CAST(value AS STRING)) AS INT) " + + val query = "SELECT CAST(IF(_1 = 1, '1', CAST(_1 AS STRING)) AS INT) " + "FROM cast_branch_controls" val (_, plan) = checkSparkAnswerAndOperator(sql(query)) assert(collect(plan) { case project: CometProjectExec => project }.nonEmpty) @@ -116,7 +119,7 @@ class CometNativeCastSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } withSQLConf(SQLConf.ANSI_ENABLED.key -> "false") { - val query = "SELECT CAST(IF(value = 1, 'bad', CAST(value AS STRING)) AS INT) " + + val query = "SELECT CAST(IF(_1 = 1, 'bad', CAST(_1 AS STRING)) AS INT) " + "FROM cast_branch_controls" val (_, plan) = checkSparkAnswerAndOperator(sql(query)) assert(collect(plan) { case project: CometProjectExec => project }.nonEmpty) diff --git a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala index 8903f43769b..6e2511ccb3c 100644 --- a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala +++ b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala @@ -452,11 +452,14 @@ abstract class CometTestBase } } - /** Checks native execution and Spark exception type, error class and SQLSTATE parity. */ + /** Checks Spark exception parity and, by default, native execution. */ protected def checkSparkError( df: DataFrame, - errorClass: String): SparkThrowable with Throwable = { - checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) + errorClass: String, + checkNative: Boolean = true): SparkThrowable with Throwable = { + if (checkNative) { + checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) + } val (sparkError, cometError) = checkSparkAnswerMaybeThrows(df) def structuredError(