如何在Scala中创建Spark SQL读取数据的Predicate列表
问题描述
我可以通过以下简单的Scala程序读取Oracle表:
val spark = SparkSession .builder .master("local[4]") .config("spark.sql.sources.partitionColumnTypeInference.enabled", false) .config("spark.executor.memory", "8g") .config("spark.executor.cores", 4) .config("spark.task.cpus", 1) .appName("Spark SQL basic example") .config("spark.some.config.option", "some-value") .getOrCreate() val jdbcDF = spark.read .format("jdbc") .option("url", "jdbc:oracle:thin:@x.x.x.x:1521:orcl") .option("dbtable", "big_table") .option("user", "test") .option("password", "123456") .load() jdbcDF.show()
但该表数据量极大,需要让每个Spark节点读取部分数据。因此必须使用哈希函数在节点间分配数据,Spark提供了Predicates来实现这一需求。我已在Python中完成该逻辑:表包含名为NUM的列,哈希函数接收列值并返回0到num_partitions之间的整数,通过生成包含日期条件和哈希匹配的Predicate列表实现数据分片读取。请问如何在Scala中实现上述Python逻辑中的Predicate列表?
Scala实现方案
以下是对应Python逻辑的Scala实现步骤:
- 定义核心变量
先声明需要用到的参数,与Python中的变量一一对应:
val numPartitions = 19 val partitionKey = "你的分区日期列名" // 替换为实际业务日期列名 val currentDate = "20240520" // 替换为实际日期字符串,格式为YYYYMMDD val hashCol = "NUM" // 用于哈希计算的目标列名 val sourceTableName = "big_table" // 要读取的Oracle表名
- 生成哈希值列表
直接生成0到numPartitions的整数序列,对应Python中的hash_values:
val hashValues = 0 to numPartitions
- 构建Predicate列表
遍历哈希值列表,拼接每个Predicate字符串,可使用Scala的字符串插值或format方法完成变量替换:
// 方式一:使用Scala字符串插值构建 val predicates = hashValues.map { hashVal => s"""to_date($partitionKey,'YYYYMMDD','nls_calendar=persian')= to_date('$currentDate','YYYYMMDD','nls_calendar=persian') |and ora_hash($hashCol, $numPartitions) = $hashVal""".stripMargin }.toArray
或者使用format方法实现:
// 方式二:使用字符串format方法 val predicateTemplate = """to_date(%s,'YYYYMMDD','nls_calendar=persian')= to_date('%s','YYYYMMDD','nls_calendar=persian') |and ora_hash(%s, %d) = %d""".stripMargin val predicates = hashValues.map { hashVal => predicateTemplate.format(partitionKey, currentDate, hashCol, numPartitions, hashVal) }.toArray
- 基于Predicate读取Oracle表
将生成的predicates数组传入JDBC读取方法,完成分片数据读取:
val connectionProps = new java.util.Properties() connectionProps.put("user", "test") connectionProps.put("password", "123456") val dataframe = spark.read .option("driver", "oracle.jdbc.driver.OracleDriver") .jdbc( url = "jdbc:oracle:thin:@x.x.x.x:1521:orcl", table = sourceTableName, predicates = predicates, connectionProperties = connectionProps ) dataframe.show()
扩展:动态获取Oracle中存在的哈希值
如果需要先查询Oracle获取表中实际存在的哈希值(而非直接生成0到numPartitions的序列),可以先通过Spark JDBC执行查询,再生成Predicate:
// 先查询表中distinct的哈希值 val hashQuery = s"SELECT DISTINCT ora_hash($hashCol, $numPartitions) AS hash FROM $sourceTableName" val hashDF = spark.read .format("jdbc") .option("url", "jdbc:oracle:thin:@x.x.x.x:1521:orcl") .option("dbtable", s"($hashQuery) tmp") .option("user", "test") .option("password", "123456") .load() // 提取哈希值到数组 val hashValues = hashDF.select("hash").as[Int].collect() // 生成Predicate列表(同之前的逻辑) val predicates = hashValues.map { hashVal => s"""to_date($partitionKey,'YYYYMMDD','nls_calendar=persian')= to_date('$currentDate','YYYYMMDD','nls_calendar=persian') |and ora_hash($hashCol, $numPartitions) = $hashVal""".stripMargin }
内容的提问来源于stack exchange,提问作者M_Gh
相关产品推荐
相关产品推荐

