Skip to content
Open
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
26 changes: 22 additions & 4 deletions spark/src/main/scala/org/apache/comet/expressions/CometCast.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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
Expand Down Expand Up @@ -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)) {
Expand All @@ -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
}
Comment thread
sunchao marked this conversation as resolved.
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)
Expand Down
67 changes: 66 additions & 1 deletion spark/src/test/scala/org/apache/comet/CometNativeCastSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 =
Expand Down
9 changes: 6 additions & 3 deletions spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading