From 96375858dd3f3110426e46377c6c6d9fb7719a5a Mon Sep 17 00:00:00 2001 From: Marko Sisovic Date: Mon, 27 Jul 2026 10:23:52 +0000 Subject: [PATCH 1/2] [SPARK-XXXXX][SQL] Truncate fractional to integral casts pushed down to JDBC sources Spark truncates toward zero when casting a fractional value to an integral type, but MySQL, Oracle, Postgres and Snowflake round half away from zero. A pushed down cast therefore returns different results than Spark evaluating it locally, silently, with no error. Wrap such casts in the dialect's truncating function via a new JdbcDialect.truncateFractionalValue hook, applied once in JDBCSQLBuilder.visitCast. The default returns the expression unchanged, so dialects whose cast already truncates (or that have no truncating function) are unaffected. spark.sql.legacy.jdbc.roundIntegralCastPushdown.enabled restores the previous behavior. --- .../apache/spark/sql/jdbc/v2/V2JDBCTest.scala | 21 ++++++++++++++++++ .../apache/spark/sql/internal/SQLConf.scala | 12 ++++++++++ .../apache/spark/sql/jdbc/JdbcDialects.scala | 22 ++++++++++++++++++- .../apache/spark/sql/jdbc/MySQLDialect.scala | 4 ++++ .../apache/spark/sql/jdbc/OracleDialect.scala | 3 +++ .../spark/sql/jdbc/PostgresDialect.scala | 3 +++ .../spark/sql/jdbc/SnowflakeDialect.scala | 3 +++ 7 files changed, 67 insertions(+), 1 deletion(-) 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..9f5e1d0108986 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() 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..97973ef083c57 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.getConf(SQLConf.LEGACY_JDBC_ROUND_INTEGRAL_CAST_PUSHDOWN)) { + 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..59a9d16019097 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,6 +50,10 @@ 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 } 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). From fc0f9a81916ab926923cca540d17e413aaee9246 Mon Sep 17 00:00:00 2001 From: Marko Sisovic Date: Mon, 27 Jul 2026 13:53:04 +0000 Subject: [PATCH 2/2] Cast integral targets to SIGNED for MySQL and add a conf accessor MySQL has no INTEGER cast target, so the base CAST(... AS INTEGER) is a syntax error there. It was never hit before because no test pushed an integral cast to MySQL. Override visitCast in MySQLSQLBuilder to render SIGNED, composing with truncateFractionalValue. Also read the new conf through a SQLConf accessor, like the other legacy JDBC confs. --- .../scala/org/apache/spark/sql/internal/SQLConf.scala | 3 +++ .../scala/org/apache/spark/sql/jdbc/JdbcDialects.scala | 2 +- .../scala/org/apache/spark/sql/jdbc/MySQLDialect.scala | 10 ++++++++++ 3 files changed, 14 insertions(+), 1 deletion(-) 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 9f5e1d0108986..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 @@ -8592,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 97973ef083c57..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 @@ -447,7 +447,7 @@ abstract class JdbcDialect extends Serializable with Logging { // Dialects for such databases override `truncateFractionalValue` to truncate the value first. val castedExpr = if (exprDataType.isInstanceOf[FractionalType] && dataType.isInstanceOf[IntegralType] && - !SQLConf.get.getConf(SQLConf.LEGACY_JDBC_ROUND_INTEGRAL_CAST_PUSHDOWN)) { + !SQLConf.get.legacyJdbcRoundIntegralCastPushdown) { truncateFractionalValue(expr) } else { expr 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 59a9d16019097..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 @@ -60,6 +60,16 @@ private case class MySQLDialect() extends JdbcDialect with SQLConfHelper with No 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 {