Spark Scala使用foreachPartition处理数据时取行首列报序列化错误如何解决
问题场景
通过Spark foreachPartition算子处理DataFrame数据批量写入数据库,代码如下:
val endDF=spark.read.parquet(path).select("pc").filter(col("pc").isNotNull); endDF.foreachPartition((partition: Iterator[Row]) => { Class.forName(driver) val con=DriverManager.getConnection(jdbcurl,user,pwd) partition.grouped(100).foreach(batch => { val st=con.createStatement() batch.foreach(row => { val pc=row.get(0).toString() val in=s"""insert tshdim (pc) values(${pc})""".stripMargin st.addBatch(in) }) st.executeLargeBatch }) con.close() })
异常信息
执行时抛出如下异常:
org.apache.spark.SparkException : Task not serializable . .
根因异常:
Java.io.NotSerializable exception:
org.apache.spark.sql.DataSet$RDDQueryExecution$ Serialization stack:
Object not serializable
(class:org.apache.spark.sql.DataSet$RDDQueryExecution$, value:
org.apache.spark.sql.DataSet$RDDQueryExecution$@jfaf )
-field(class:org.apache.spark.sql.DataSet, name:RDDQueryExecutionModule, type:
org.apache.spark.sql.DataSet$RDDQueryExecution$)
-object(class:org.apache.spark.sql.DataSet,[pc:String])
根因排查
该异常属于Spark常见的闭包序列化问题:
Spark需要将foreachPartition算子传入的函数(闭包)序列化后分发到各Executor节点执行,如果闭包引用了Driver端不可序列化的对象,就会触发该错误。本次问题的触发原因是:
- 闭包内直接引用了Driver端定义的
driver、jdbcurl、user、pwd等变量,闭包捕获时连带引用了关联的DataSet不可序列化对象,导致序列化失败。 - 代码中
foreachpartition拼写错误(应为foreachPartition,P大写),虽不是本次异常的直接原因,但会导致算子调用失败。
解决办法
修正方案
- 所有需要传入Executor的参数,要么定义为独立的序列化常量,要么使用广播变量传递,避免闭包捕获不必要的外部对象。
- 数据库操作改用预编译
PreparedStatement,既避免SQL注入风险,也提升批处理性能。 - 增加连接、Statement的资源释放逻辑,避免异常情况下资源泄漏。
修正后代码
// 提前将配置参数定义为局部常量,避免闭包关联DataSet对象 val jdbcDriver = driver val url = jdbcurl val username = user val password = pwd val endDF=spark.read.parquet(path).select("pc").filter(col("pc").isNotNull) endDF.foreachPartition((partition: Iterator[Row]) => { // 所有数据库连接相关操作均在闭包内部执行,不依赖外部不可序列化对象 Class.forName(jdbcDriver) var con: Connection = null var pst: PreparedStatement = null try { con = DriverManager.getConnection(url, username, password) // 开启手动提交,提升批处理性能 con.setAutoCommit(false) val insertSql = "insert into tshdim (pc) values (?)" pst = con.prepareStatement(insertSql) partition.grouped(100).foreach(batch => { batch.foreach(row => { val pc = row.getString(0) pst.setString(1, pc) pst.addBatch() }) pst.executeLargeBatch() con.commit() pst.clearBatch() }) } catch { case e: Exception => if (con != null) con.rollback() throw e } finally { // 资源释放 if (pst != null) pst.close() if (con != null) con.close() } })
内容的提问来源于stack exchange,提问作者nirmal

