Spark多分区读MySQL后Join报SQL语法错的原理咨询及建议
问题场景
使用Spark多分区并行读取两个MySQL表构建DataFrame后执行Join操作,触发MySQL语法错误,报错信息如下:
com.mysql.jdbc.exceptions.jdbc4.MySQLSyntaxErrorException: You have an error in your SQL syntax; check the manual that corresponds to your MySQL server version for the right syntax to use near 'limit 19796,4949)' at line 1 at sun.reflect.NativeConstructorAccessorImpl.newInstance0(Native Method) at sun.reflect.NativeConstructorAccessorImpl.newInstance(NativeConstructorAccessorImpl.java:62) at sun.reflect.DelegatingConstructorAccessorImpl.newInstance(DelegatingConstructorAccessorImpl.java:45) at java.lang.reflect.Constructor.newInstance(Constructor.java:423) at com.mysql.jdbc.Util.handleNewInstance(Util.java:425) at com.mysql.jdbc.Util.getInstance(Util.java:408) at com.mysql.jdbc.SQLError.createSQLException(SQLError.java:944) at com.mysql.jdbc.MysqlIO.checkErrorPacket(MysqlIO.java:3978) at com.mysql.jdbc.MysqlIO.checkErrorPacket(MysqlIO.java:3914) at com.mysql.jdbc.MysqlIO.sendCommand(MysqlIO.java:2530) at com.mysql.jdbc.MysqlIO.sqlQueryDirect(MysqlIO.java:2683) at com.mysql.jdbc.ConnectionImpl.execSQL(ConnectionImpl.java:2495) at com.mysql.jdbc.PreparedStatement.executeInternal(PreparedStatement.java:1903) at com.mysql.jdbc.PreparedStatement.executeQuery(PreparedStatement.java:2011) at org.apache.spark.sql.execution.datasources.jdbc.JDBCRDD.compute(JDBCRDD.scala:358) at org.apache.spark.rdd.RDD.computeOrReadCheckpoint(RDD.scala:373) at org.apache.spark.rdd.RDD.iterator(RDD.scala:337) at org.apache.spark.rdd.MapPartitionsRDD.compute(MapPartitionsRDD.scala:52) at org.apache.spark.rdd.RDD.computeOrReadCheckpoint(RDD.scala:373) at org.apache.spark.rdd.RDD.iterator(RDD.scala:337) at org.apache.spark.rdd.MapPartitionsRDD.compute(MapPartitionsRDD.scala:52) at org.apache.spark.rdd.RDD.computeOrReadCheckpoint(RDD.scala:373) at org.apache.spark.rdd.RDD.iterator(RDD.scala:337) at org.apache.spark.shuffle.ShuffleWriteProcessor.write(ShuffleWriteProcessor.scala:59) at org.apache.spark.scheduler.ShuffleMapTask.runTask(ShuffleMapTask.scala:99) at org.apache.spark.scheduler.ShuffleMapTask.runTask(ShuffleMapTask.scala:52) at org.apache.spark.scheduler.Task.run(Task.scala:131) at org.apache.spark.executor.Executor$TaskRunner.$anonfun$run$3(Executor.scala:506) at org.apache.spark.util.Utils$.tryWithSafeFinally(Utils.scala:1462) at org.apache.spark.executor.Executor$TaskRunner.run(Executor.scala:509) at java.util.concurrent.ThreadPoolExecutor.runWorker(ThreadPoolExecutor.java:1149) at java.util.concurrent.ThreadPoolExecutor$Worker.run(ThreadPoolExecutor.java:624) at java.lang.Thread.run(Thread.java:748)
相关代码
主程序代码
object TestJoin { def main(args: Array[String]): Unit = { if (System.getProperty("os.name").toLowerCase().contains("win")) { InitSparkEnv.init(args(0), "local[*]") } else { InitSparkEnv.initNotSupportHive(args(0)) } val ssc = InitSparkEnv.getSparkSession val tableName_goods = "data_center_cdm.dwd_cob_mc_enterprise_goods" val tableName_goods_attr = "data_center_cdm.dwd_cob_mc_enterprise_goods_attr" ConnUtil.initMysql() val prop = ConnUtil.getProp val readFromMysql = new ReadFromMysql(ssc, prop) // 5 partitions read in parallel val goods = readFromMysql.getDataByPage(prop.getProperty("url"), tableName_goods, 5) val goods_attr = readFromMysql.getDataByPage(prop.getProperty("url"), tableName_goods_attr, 5) val result = goods.join(goods_attr, "goods_md5") result.show(false) ssc.stop() } }
多分区读取工具类代码
class ReadFromMysql(ssc: SparkSession) extends Serializable { private val NUM_PARTITIONS = 20 private val MAX_FETCH_SIZE = 100 /** * * @param url url * @param tableName tableName * @param pageNum PartitionNum */ def getDataByPage(url: String, tableName: String, pageNum: Int): DataFrame = { // 查询该表的数量级 val query_sql = s"(select count(*) from ${tableName}) tbl" val tableRows = getData(url, query_sql) val tableNumRecords = tableRows.head().get(0).asInstanceOf[Long] if (tableNumRecords <= MAX_FETCH_SIZE) { getData(url, tableName) } else { val usePartNum = if (pageNum > NUM_PARTITIONS) { logger.warn("user defined num-partitions:" + pageNum + ",but max partitions is " + NUM_PARTITIONS) min(pageNum, NUM_PARTITIONS) } else { logger.info("user defined num-partitions:" + pageNum + "") pageNum } val predicates: ArrayBuffer[String] = ArrayBuffer[String]() val setupNums = tableNumRecords / usePartNum val remainder = tableNumRecords - (setupNums * usePartNum) logger.info("The number of data pulled each time is " + setupNums) for (num <- 0 until usePartNum) { predicates += "1=1 limit " + setupNums * num + "," + setupNums } predicates += "1=1 limit " + (setupNums * usePartNum) + "," + remainder ssc.read.jdbc(url, tableName, predicates.toArray, prop) } } }
临时解决方法
- 缓存DataFrame:对读取后的DataFrame调用
cache()方法,可避免报错,修改后的代码片段:
// 5 partitions read in parallel val goods = readFromMysql.getDataByPage(prop.getProperty("url"), tableName_goods, 5) val goods_attr = readFromMysql.getDataByPage(prop.getProperty("url"), tableName_goods_attr, 5) goods.cache() goods_attr.cache() val result = goods.join(goods_attr, "goods_md5") result.show(false) ssc.stop()
- 单分区读取:使用单分区读取数据也无此问题,单分区读取方法代码:
/** * * @param url url * @param tableName tableMane * @return */ def getData(url: String, tableName: String): DataFrame = { logger.warn(s"this method read data just by one precess,start reading table [${tableName}] >>>>>>>>>>>>>>>>>>>>>>>>") ssc.read.jdbc(url, tableName, prop) }
疑问
不清楚缓存解决问题的底层原理,希望了解原因并获得相关优化建议。
一、缓存解决问题的底层原理
1. Spark的懒执行特性
Spark的DataFrame采用懒执行机制,在未调用show()、count()等action类算子前,不会实际执行数据读取和计算操作。直接执行Join时,Spark会尝试将整个查询计划下推到MySQL,即谓词下推。
2. 多分区读取的SQL构造问题
当前多分区读取逻辑通过predicates数组传递limit offset, rows条件,Spark执行Join时,会将Join逻辑与分区的limit条件组合后下推到MySQL,生成的SQL类似:
SELECT * FROM table WHERE (1=1 limit 19796,4949) AND goods_md5 IN (...)
而MySQL语法不允许limit作为WHERE条件的一部分,因此触发语法错误。
3. 缓存的作用
调用cache()后,Spark会提前触发数据读取的action操作,将数据加载到内存/磁盘中。此时Join操作基于内存中的数据执行,不再将查询计划下推到MySQL,自然不会构造出错误的SQL语句,从而避免报错。
二、优化建议
1. 修正多分区读取的分区策略
避免用limit offset作为分区条件,该方式不仅引发语法问题,还会因偏移量增大导致MySQL查询性能下降。建议基于表中的有序主键或唯一键做范围分区:
// 假设表中有id列作为主键 val minId = // 查询表的最小id val maxId = // 查询表的最大id val step = (maxId - minId) / usePartNum val predicates = (0 until usePartNum).map(i => { val start = minId + i * step val end = if (i == usePartNum -1) maxId else start + step s"id >= $start AND id <= $end" }).toArray ssc.read.jdbc(url, tableName, predicates, prop)
这种方式每个分区执行范围查询,SQL语法合法,且MySQL可利用主键索引提升查询效率。
2. 关闭不必要的谓词下推
若必须使用limit分区方式,可关闭JDBC数据源的谓词下推:
val df = ssc.read .option("pushDownPredicate", "false") .jdbc(url, tableName, predicates, prop)
这样Spark会先将所有数据读取到节点后再执行计算,避免语法错误,但会增加数据传输量,需根据场景权衡。
3. 合理设置分区数
分区数不宜超过MySQL的max_connections配置,否则会引发连接失败。建议结合数据量、Spark executor资源和MySQL可用连接数调整分区数。
4. 预读取数据到中间存储
若需多次使用这些表的数据,可提前将数据读取到HDFS、Hive等分布式存储中,后续直接从这些存储读取数据执行Join,避免重复从MySQL读取,提升整体性能。
内容的提问来源于stack exchange,提问作者Mr.Guo

