Spark3.1.2(Scala2.12)下RDD[String]转DataFrame报错如何解决
问题根源
你报错的直接原因是 spark.createDataFrame 传入 StructType 类型的Schema时,要求第一个参数必须是 RDD[org.apache.spark.sql.Row] 类型,但你当前传入的是 RDD[String],类型不匹配所以重载方法找不到对应的实现。
另外你代码里还有一个隐藏的逻辑错误:read_data 方法里把所有查询结果的所有字段直接拼接成了一个完整字符串,完全丢失了行和字段的边界,就算类型改对也没法正常解析成DataFrame。
修复步骤
- 调整
read_data方法的返回值,不要把结果拼接成字符串,直接返回每一行对应的Row对象,同时正确处理字段类型、补充资源关闭逻辑避免连接泄漏:
// 先导入必要的依赖类 import org.apache.spark.sql.Row import java.sql.{Connection, Statement, ResultSet} def read_data(group_id: Int): Iterator[Row] = { val table_name = "table" val col_name = "col" val query = s""" select f1,f2,f3,f4,f5,f6,f7,f8 | from $table_name | where MOD(TO_NUMBER(substr($col_name, -LEAST(2, LENGTH($col_name)))),$num_node)=$group_id """.stripMargin val oracleUser = "ORCL" val oraclePassword = "XXXXXXXX" val oracleURL = "jdbc:oracle:thin:@//X.X.X.X:1521/ORCLDB" val ods = new OracleDataSource() ods.setUser(oracleUser) ods.setURL(oracleURL) ods.setPassword(oraclePassword) var con: Connection = null var statement: Statement = null var resultSet: ResultSet = null try { con = ods.getConnection() statement = con.createStatement() statement.setFetchSize(1000) resultSet = statement.executeQuery(query) Iterator.continually(resultSet) .takeWhile(_.next) .map(r => Row( r.getString(1), r.getString(2), r.getString(3), r.getString(4), r.getString(5), r.getString(6), r.getInt(7), // 匹配schema中f7的IntegerType r.getInt(8) // 匹配schema中f8的IntegerType )) } finally { if (resultSet != null) resultSet.close() if (statement != null) statement.close() if (con != null) con.close() } }
- 调整RDD生成逻辑,把
map改成flatMap展开每个分组返回的行迭代器,得到RDD[Row]:
val rdd = sc.parallelize(group_list,num_node) .flatMap(read_data)
- 修改转换方法的入参类型即可正常调用
createDataFrame:
// 方法名不要用udf,避免和Spark原生UDF概念混淆 def rddToDf(rdd: RDD[Row]): DataFrame = { val spark = SparkSession.builder.getOrCreate() val schema = new StructType() .add(StructField("f1", StringType, true)) .add(StructField("f2", StringType, true)) .add(StructField("f3", StringType, true)) .add(StructField("f4", StringType, true)) .add(StructField("f5", StringType, true)) .add(StructField("f6", StringType, true)) .add(StructField("f7", IntegerType, true)) .add(StructField("f8", IntegerType, true)) spark.createDataFrame(rdd, schema) }
其他优化建议
- 你当前的
group_list = Seq.range(1,2,1)只有1一个值,Oracle的MOD函数返回结果范围是0~num_node-1,会漏掉分组id为0的数据,建议改成Seq.range(0, num_node) - 不要在Spark转换算子中直接执行打印逻辑,
rdd.map(x => println(x)).collect()会把打印分发到executor执行,driver端看不到输出,改成rdd.collect().foreach(println)即可。
内容的提问来源于stack exchange,提问作者M_Gh
相关产品推荐
相关产品推荐

