diff --git a/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/V2JDBCTest.scala b/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/V2JDBCTest.scala index 6f2a9e97ff005..51c11fbf0cba6 100644 --- a/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/V2JDBCTest.scala +++ b/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/V2JDBCTest.scala @@ -1362,6 +1362,27 @@ private[v2] trait V2JDBCTest } } + test("fractional to integral cast pushed down truncates toward zero like Spark") { + val tbl = s"$catalogName.integral_cast" + withTable(tbl) { + // label is a string so the result type is the same across dialects (a numeric column + // comes back as BigDecimal on some engines such as Oracle). + sql(s"CREATE TABLE $tbl (double_col DOUBLE, decimal_col DECIMAL(10, 2), label VARCHAR(8))") + sql(s"INSERT INTO $tbl VALUES (1.5, 1.5, 'a'), (2.5, 2.5, 'b'), (-1.5, -1.5, 'c')") + + // Spark truncates toward zero, so 1.5, 2.5 and -1.5 become 1, 2 and -1. A database that + // rounds half away from zero would return 2, 3 and -2 instead. + Seq("double_col", "decimal_col").foreach { col => + val projected = sql(s"SELECT CAST($col AS INT) FROM $tbl ORDER BY $col") + assert(projected.collect().map(_.getInt(0)) === Array(-1, 1, 2)) + + val filtered = sql(s"SELECT label FROM $tbl WHERE CAST($col AS INT) = 2") + checkFilterPushed(filtered) + assert(filtered.collect().map(_.getString(0)) === Array("b")) + } + } + } + test("SPARK-52262: FAILED_JDBC.TABLE_EXISTS not thrown on connection error") { val invalidTableName = s"$catalogName.invalid" val originalUrl = spark.conf.get(s"spark.sql.catalog.$catalogName.url") diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala index 77c69494a15a5..373fcf2f9ccc3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala @@ -6569,6 +6569,18 @@ object SQLConf { .booleanConf .createWithDefault(false) + val LEGACY_JDBC_ROUND_INTEGRAL_CAST_PUSHDOWN = + buildConf("spark.sql.legacy.jdbc.roundIntegralCastPushdown.enabled") + .internal() + .doc("When true, a cast from a fractional type to an integral type that is pushed down to " + + "a JDBC data source is not wrapped in the dialect's truncating function, so the result " + + "follows the database's own cast semantics (the legacy behavior). Databases that round " + + "such casts then return different results than Spark, which truncates toward zero.") + .version("4.3.0") + .withBindingPolicy(ConfigBindingPolicy.SESSION) + .booleanConf + .createWithDefault(false) + val LEGACY_JDBC_TIME_MAPPING_ENABLED = buildConf("spark.sql.legacy.jdbc.timeMapping.enabled") .internal() @@ -8580,6 +8592,9 @@ class SQLConf extends Serializable with Logging with SqlApiConf { def legacyMsSqlServerDatetimeOffsetMappingEnabled: Boolean = getConf(LEGACY_MSSQLSERVER_DATETIMEOFFSET_MAPPING_ENABLED) + def legacyJdbcRoundIntegralCastPushdown: Boolean = + getConf(LEGACY_JDBC_ROUND_INTEGRAL_CAST_PUSHDOWN) + def legacyMySqlBitArrayMappingEnabled: Boolean = getConf(LEGACY_MYSQL_BIT_ARRAY_MAPPING_ENABLED) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/JdbcDialects.scala b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/JdbcDialects.scala index bfc01f406dd97..e493aab9ae444 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/JdbcDialects.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/JdbcDialects.scala @@ -442,7 +442,17 @@ abstract class JdbcDialect extends Serializable with Logging { override def visitCast(expr: String, exprDataType: DataType, dataType: DataType): String = { val databaseTypeDefinition = getJDBCType(dataType).map(_.databaseTypeDefinition).getOrElse(dataType.typeName) - s"CAST($expr AS $databaseTypeDefinition)" + // Spark truncates toward zero when casting a fractional value to an integral type, but many + // databases round instead, so a pushed down cast can return a different result than Spark. + // Dialects for such databases override `truncateFractionalValue` to truncate the value first. + val castedExpr = if (exprDataType.isInstanceOf[FractionalType] && + dataType.isInstanceOf[IntegralType] && + !SQLConf.get.legacyJdbcRoundIntegralCastPushdown) { + truncateFractionalValue(expr) + } else { + expr + } + s"CAST($castedExpr AS $databaseTypeDefinition)" } override def visitSQLFunction(funcName: String, inputs: Array[Expression]): String = { @@ -525,6 +535,16 @@ abstract class JdbcDialect extends Serializable with Logging { } } + /** + * Truncates `expr` toward zero, so that a cast from a fractional type to an integral type that + * is pushed down matches Spark, which truncates rather than rounds. Dialects for databases that + * round such casts override this with the database's truncating function. The default returns + * `expr` unchanged, for databases whose cast already truncates like Spark. + * @param expr The SQL string of the value being cast. + * @return The SQL string to cast instead of `expr`. + */ + def truncateFractionalValue(expr: String): String = expr + /** * Returns whether the database supports function. * @param funcName Upper-cased function name diff --git a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/MySQLDialect.scala b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/MySQLDialect.scala index b301c0c0bd5bc..479f1558fccc9 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/MySQLDialect.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/MySQLDialect.scala @@ -50,12 +50,26 @@ private case class MySQLDialect() extends JdbcDialect with SQLConfHelper with No override def isSupportedFunction(funcName: String): Boolean = supportedFunctions.contains(funcName) + // MySQL rounds half away from zero when casting to an integral type. It has no single argument + // TRUNCATE, so the number of decimal places to keep is passed explicitly. + override def truncateFractionalValue(expr: String): String = s"TRUNCATE($expr, 0)" + override def isObjectNotFoundException(e: SQLException): Boolean = { e.getErrorCode == 1146 } class MySQLSQLBuilder extends JDBCSQLBuilder { + override def visitCast(expr: String, exprDataType: DataType, dataType: DataType): String = + dataType match { + // MySQL does not support casting to SHORT, INT or BIGINT, it uses SIGNED instead. + case _: IntegralType if exprDataType.isInstanceOf[FractionalType] && + !conf.legacyJdbcRoundIntegralCastPushdown => + s"CAST(${truncateFractionalValue(expr)} AS SIGNED)" + case _: IntegralType => s"CAST($expr AS SIGNED)" + case _ => super.visitCast(expr, exprDataType, dataType) + } + override def visitExtract(extract: Extract): String = { val field = extract.field field match { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/OracleDialect.scala b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/OracleDialect.scala index f5eb1fa6d7a06..a9499c25f2cfe 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/OracleDialect.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/OracleDialect.scala @@ -49,6 +49,9 @@ private case class OracleDialect() extends JdbcDialect with SQLConfHelper with N override def isSupportedFunction(funcName: String): Boolean = supportedFunctions.contains(funcName) + // Oracle rounds half away from zero when casting to an integral type. + override def truncateFractionalValue(expr: String): String = s"TRUNC($expr)" + override def isObjectNotFoundException(e: SQLException): Boolean = { e.getMessage.contains("ORA-00942") || e.getMessage.contains("ORA-39165") diff --git a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/PostgresDialect.scala b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/PostgresDialect.scala index 9fc69594932ba..8bad6fdcc3b53 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/PostgresDialect.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/PostgresDialect.scala @@ -53,6 +53,9 @@ private case class PostgresDialect() override def isSupportedFunction(funcName: String): Boolean = supportedFunctions.contains(funcName) + // Postgres rounds half away from zero when casting to an integral type. + override def truncateFractionalValue(expr: String): String = s"TRUNC($expr)" + override def isObjectNotFoundException(e: SQLException): Boolean = { e.getSQLState == "42P01" || e.getSQLState == "3F000" || diff --git a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/SnowflakeDialect.scala b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/SnowflakeDialect.scala index 1c88a554863ab..a68f2c7780e44 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/jdbc/SnowflakeDialect.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/jdbc/SnowflakeDialect.scala @@ -31,6 +31,9 @@ private case class SnowflakeDialect() extends JdbcDialect with NoLegacyJDBCError e.getSQLState == "002003" } + // Snowflake rounds half away from zero when casting to an integral type. + override def truncateFractionalValue(expr: String): String = s"TRUNC($expr)" + override def getJDBCType(dt: DataType): Option[JdbcType] = dt match { case BooleanType => // By default, BOOLEAN is mapped to BIT(1).