diff --git a/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/MySQLIntegrationSuite.scala b/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/MySQLIntegrationSuite.scala index c2714587e2d2..fc2cad1b7098 100644 --- a/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/MySQLIntegrationSuite.scala +++ b/connector/docker-integration-tests/src/test/scala/org/apache/spark/sql/jdbc/v2/MySQLIntegrationSuite.scala @@ -20,7 +20,7 @@ package org.apache.spark.sql.jdbc.v2 import java.sql.{Connection, SQLFeatureNotSupportedException} import org.apache.spark.{SparkConf, SparkSQLFeatureNotSupportedException} -import org.apache.spark.sql.AnalysisException +import org.apache.spark.sql.{AnalysisException, Row} import org.apache.spark.sql.execution.datasources.v2.jdbc.JDBCTableCatalog import org.apache.spark.sql.jdbc.MySQLDatabaseOnDocker import org.apache.spark.sql.types._ @@ -313,6 +313,15 @@ class MySQLIntegrationSuite extends DockerJDBCIntegrationV2Suite with V2JDBCTest assert(rows10(0).getString(0) === "amy") assert(rows10(1).getString(0) === "alex") } + + test("do not push down casts to double") { + val df = sql( + s"SELECT name FROM $catalogName.employee " + + "WHERE CAST(salary AS DOUBLE) > 10000.5") + + checkFilterPushed(df, pushed = false) + checkAnswer(df, Seq(Row("alex"), Row("jen"))) + } } /** 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 60cce5babe4c..ed533795329f 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 @@ -118,6 +118,16 @@ private case class MySQLDialect() extends JdbcDialect with SQLConfHelper with No } else { super.visitAggregateFunction(funcName, isDistinct, inputs) } + + override def visitCast(expr: String, exprDataType: DataType, dataType: DataType): String = { + dataType match { + case DoubleType => + // The common JDBC mapping is DOUBLE PRECISION, which is not a portable CAST target for + // MySQL-compatible databases. In particular, MariaDB rejects it with a syntax error. + throw new UnsupportedOperationException("Cannot cast to double type") + case _ => super.visitCast(expr, exprDataType, dataType) + } + } } override def compileExpression(expr: Expression): Option[String] = { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/jdbc/JDBCSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/jdbc/JDBCSuite.scala index b654deb12ad4..0213d18922b7 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/jdbc/JDBCSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/jdbc/JDBCSuite.scala @@ -38,7 +38,7 @@ import org.apache.spark.sql.catalyst.parser.CatalystSqlParser import org.apache.spark.sql.catalyst.plans.logical.ShowCreateTable import org.apache.spark.sql.catalyst.util.{CaseInsensitiveMap, CharVarcharUtils, DateTimeTestUtils} import org.apache.spark.sql.connector.catalog.Identifier -import org.apache.spark.sql.connector.expressions.{Expression => V2Expression, FieldReference, GeneralScalarExpression, LiteralValue} +import org.apache.spark.sql.connector.expressions.{Cast => V2Cast, Expression => V2Expression, FieldReference, GeneralScalarExpression, LiteralValue} import org.apache.spark.sql.connector.expressions.filter.{AlwaysFalse, AlwaysTrue, Predicate} import org.apache.spark.sql.execution.{DataSourceScanExec, ExtendedMode, ProjectExec} import org.apache.spark.sql.execution.command.{ExplainCommand, ShowCreateTableCommand} @@ -1454,6 +1454,13 @@ class JDBCSuite extends SharedSparkSession { assert(mySqlDialect.getJDBCType(FloatType).map(_.databaseTypeDefinition).get == "FLOAT") } + test("MySQL blocks casts to double") { + val dialect = MySQLDialect() + val cast = new V2Cast(FieldReference("value"), IntegerType, DoubleType) + + assert(dialect.compileExpression(cast).isEmpty) + } + test("PostgresDialect type mapping") { val Postgres = JdbcDialects.get("jdbc:postgresql://127.0.0.1/db") val md = new MetadataBuilder().putLong("scale", 0).putBoolean("isTimestampNTZ", false)