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 cce753a68fff5..95acc05ba3e8e 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 @@ -52,7 +52,7 @@ import org.apache.spark.sql.catalyst.util.DateTimeUtils import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME import org.apache.spark.sql.connector.catalog.PathElement.PathRef import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors} -import org.apache.spark.sql.types.{AtomicType, TimestampNTZType, TimestampType} +import org.apache.spark.sql.types.{AtomicType, DecimalType, TimestampNTZType, TimestampType} import org.apache.spark.storage.{StorageLevel, StorageLevelMapper} import org.apache.spark.unsafe.array.ByteArrayMethods import org.apache.spark.util.{HadoopFSUtils, Utils, VersionUtils} @@ -6743,6 +6743,18 @@ object SQLConf { .booleanConf .createWithDefault(false) + val JDBC_ORACLE_NUMBER_DEFAULT_SCALE = + buildConf("spark.sql.jdbc.oracle.numberDefaultScale") + .doc("Default scale for Oracle NUMBER columns that have no explicit precision/scale. " + + "Values with more fractional digits than this scale are silently rounded. " + + "Can be overridden per source via the oracle.numberDefaultScale JDBC read option.") + .version("5.0.0") + .withBindingPolicy(ConfigBindingPolicy.SESSION) + .intConf + .checkValue(s => s >= 0 && s <= DecimalType.MAX_SCALE, + s"The scale must be between 0 and ${DecimalType.MAX_SCALE}, inclusive.") + .createWithDefault(10) + val LEGACY_DB2_TIMESTAMP_MAPPING_ENABLED = buildConf("spark.sql.legacy.db2.numericMapping.enabled") .internal() @@ -8836,6 +8848,9 @@ class SQLConf extends Serializable with Logging with SqlApiConf { def legacyOracleTimestampMappingEnabled: Boolean = getConf(LEGACY_ORACLE_TIMESTAMP_MAPPING_ENABLED) + def jdbcOracleNumberDefaultScale: Int = + getConf(JDBC_ORACLE_NUMBER_DEFAULT_SCALE) + def legacyDB2numericMappingEnabled: Boolean = getConf(LEGACY_DB2_TIMESTAMP_MAPPING_ENABLED) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCOptions.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCOptions.scala index 7188c2b8b2e8e..8008af095b123 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCOptions.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCOptions.scala @@ -29,7 +29,7 @@ import org.apache.spark.internal.Logging import org.apache.spark.sql.catalyst.util.CaseInsensitiveMap import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.TimestampNTZType +import org.apache.spark.sql.types.{DecimalType, TimestampNTZType} import org.apache.spark.util.Utils /** @@ -262,6 +262,14 @@ class JDBCOptions( s"$value " }).getOrElse("") + val oracleNumberDefaultScale = parameters.get(JDBC_ORACLE_NUMBER_DEFAULT_SCALE).map { v => + val scale = v.toInt + require(scale >= 0 && scale <= DecimalType.MAX_SCALE, + s"Invalid value `$v` for option `$JDBC_ORACLE_NUMBER_DEFAULT_SCALE`." + + s" The scale must be between 0 and ${DecimalType.MAX_SCALE}, inclusive.") + scale + } + override def hashCode: Int = this.parameters.hashCode() override def equals(other: Any): Boolean = other match { @@ -367,4 +375,5 @@ object JDBCOptions { val JDBC_PREPARE_QUERY = newOption("prepareQuery") val JDBC_PREFER_TIMESTAMP_NTZ = newOption("preferTimestampNTZ") val JDBC_HINT_STRING = newOption("hint") + val JDBC_ORACLE_NUMBER_DEFAULT_SCALE = newOption("oracle.numberDefaultScale") } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCRDD.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCRDD.scala index 2989c0975143f..64dc6f3c065ae 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCRDD.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JDBCRDD.scala @@ -106,7 +106,8 @@ object JDBCRDD extends Logging { statement.setQueryTimeout(options.queryTimeout) Using.resource(statement.executeQuery()) { rs => JdbcUtils.getSchema(conn, rs, dialect, alwaysNullable = true, - isTimestampNTZ = options.preferTimestampNTZ) + isTimestampNTZ = options.preferTimestampNTZ, + oracleNumberDefaultScale = options.oracleNumberDefaultScale) } } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JdbcUtils.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JdbcUtils.scala index ca44e8b710b13..4cc518972727d 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JdbcUtils.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/jdbc/JdbcUtils.scala @@ -262,7 +262,8 @@ object JdbcUtils extends Logging with SQLConfHelper { try { statement.setQueryTimeout(options.queryTimeout) Some(getSchema(conn, statement.executeQuery(), dialect, - isTimestampNTZ = options.preferTimestampNTZ)) + isTimestampNTZ = options.preferTimestampNTZ, + oracleNumberDefaultScale = options.oracleNumberDefaultScale)) } catch { case _: SQLException => None } finally { @@ -286,7 +287,8 @@ object JdbcUtils extends Logging with SQLConfHelper { resultSet: ResultSet, dialect: JdbcDialect, alwaysNullable: Boolean = false, - isTimestampNTZ: Boolean = false): StructType = { + isTimestampNTZ: Boolean = false, + oracleNumberDefaultScale: Option[Int] = None): StructType = { val rsmd = resultSet.getMetaData val ncols = rsmd.getColumnCount val fields = new Array[StructField](ncols) @@ -328,6 +330,7 @@ object JdbcUtils extends Logging with SQLConfHelper { metadata.putBoolean("isTimestampNTZ", isTimestampNTZ) metadata.putLong("scale", fieldScale) metadata.putString("jdbcClientType", typeName) + oracleNumberDefaultScale.foreach(s => metadata.putLong("numberDefaultScale", s)) dialect.updateExtraColumnMeta(conn, rsmd, i + 1, metadata) val columnType = 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 b20faf52300ea..d78cc74959d4c 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 @@ -135,6 +135,10 @@ abstract class JdbcDialect extends Serializable with Logging { *