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..a3193931f06 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} @@ -41,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 @@ -98,9 +104,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 +127,19 @@ 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(_) => + // 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) 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/CometNativeCastSuite.scala b/spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala index 4db950e73c9..73949a2c320 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,70 @@ 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", + checkNative = false) + } + } + } + } + + test("successful literal casts and non-ANSI invalid casts retain native projection") { + 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(_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) + assert( + !plan.exists(_.getTagValue(CometExplainInfo.FALLBACK_REASONS) + .exists(_.contains(CometCast.literalCastConditionalEvalReason)))) + } + } + withSQLConf(SQLConf.ANSI_ENABLED.key -> "false") { + 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) + 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/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(