Scala Spark递归展平嵌套DataFrame为多DataFrame的问题
Scala Spark处理多层嵌套Kafka数组数据,拆分生成独立DataFrame
核心思路
要解决多层嵌套数组的拆分问题,需要递归遍历Schema定位所有层级的数组字段,递归展平所有嵌套结构,然后对每个数组字段执行explode并保留上层关联字段,最终为每个嵌套数组生成完全展平的独立DataFrame。
代码实现
1. 递归获取所有数组字段(含完整路径)
该函数会遍历Schema的所有层级,找出所有数组类型字段,并记录其完整路径(如Shipment.Shipment_Level2):
import org.apache.spark.sql.types._ def getAllArrayFields(schema: StructType, parentPath: String = ""): List[(String, ArrayType)] = { schema.fields.flatMap { field => val currentPath = if (parentPath.isEmpty) field.name else s"$parentPath.${field.name}" field.dataType match { case arrayType: ArrayType => // 记录当前数组字段,再递归遍历数组内部的结构 (currentPath, arrayType) :: getAllArrayFields(arrayType.elementType.asInstanceOf[StructType], currentPath) case structType: StructType => // 递归遍历嵌套结构体 getAllArrayFields(structType, currentPath) case _ => Nil } }.toList }
2. 递归展平嵌套Schema
将所有嵌套的结构体字段展平为一级列,列名用下划线替换路径中的点(如Shipment.Shipment_Level2.id转为Shipment_Shipment_Level2_id):
import org.apache.spark.sql.Column import org.apache.spark.sql.functions.col def flattenSchema(schema: StructType, prefix: String = ""): List[Column] = { schema.fields.flatMap { field => val colName = if (prefix.isEmpty) field.name else s"$prefix.${field.name}" field.dataType match { case structType: StructType => flattenSchema(structType, colName) case _ => List(col(colName).alias(colName.replace(".", "_"))) } }.toList }
3. 批量处理所有嵌套数组,生成独立DataFrame
对每个数组字段执行explode,保留上层关联字段,并展平数组内部的嵌套结构:
import org.apache.spark.sql.{DataFrame, SparkSession} import org.apache.spark.sql.functions.explode def processNestedArrays(baseDF: DataFrame, spark: SparkSession): Map[String, DataFrame] = { val arrayFields = getAllArrayFields(baseDF.schema) arrayFields.map { case (arrayPath, arrayType) => // 提取上层所有非数组字段作为关联键 val parentFieldNames = arrayPath.split("\\.").dropRight(1) val parentSchema = if (parentFieldNames.isEmpty) baseDF.schema else { parentFieldNames.foldLeft(baseDF.schema) { (schema, fieldName) => schema(fieldName).dataType.asInstanceOf[StructType] } } val parentCols = flattenSchema(parentSchema) // 展平数组内部的嵌套结构 val arrayElementSchema = arrayType.elementType.asInstanceOf[StructType] val flattenedArrayCols = flattenSchema(arrayElementSchema, arrayPath) // 执行explode并组合关联字段与数组字段 val explodedDF = baseDF .select(parentCols ++ List(explode(col(arrayPath)).alias("exploded_element"))) .select(parentCols ++ flattenSchema(arrayElementSchema, "exploded_element")) // 生成表名(用下划线替换路径中的点) val tableName = arrayPath.replace(".", "_") (tableName, explodedDF) }.toMap }
4. 完整业务流程示例(消费Kafka+处理+写入表)
// 初始化SparkSession val spark = SparkSession.builder() .appName("KafkaNestedArrayProcessor") .getOrCreate() // 定义Kafka Topic的Schema(替换为你的实际Schema) val yourSchema = StructType(Seq( StructField("OrderId", StringType), StructField("Shipment", StructType(Seq( StructField("Shipment_Level2", ArrayType(StructType(Seq( StructField("Level2Id", StringType), StructField("Level2RefNumbs", ArrayType(StructType(Seq( StructField("RefNumbId", StringType), StructField("RefValue", StringType) )))) )))) ))) ) // 消费Kafka并解析JSON数据 val kafkaDF = spark.readStream .format("kafka") .option("kafka.bootstrap.servers", "your-bootstrap-servers:9092") .option("subscribe", "your-target-topic") .load() .select(from_json(col("value").cast("string"), yourSchema).alias("data")) .select("data.*") // 处理所有嵌套数组,得到每个数组对应的DataFrame val tableDFs = processNestedArrays(kafkaDF, spark) // 将每个DataFrame写入对应表(以JDBC为例,流处理用foreachBatch) tableDFs.foreach { case (tableName, df) => df.writeStream .foreachBatch { (batchDF: DataFrame, _: Long) => batchDF.write .mode("append") .format("jdbc") .option("url", "jdbc:mysql://your-db-host:3306/your-db") .option("dbtable", tableName) .option("user", "db-user") .option("password", "db-password") .save() } .start() } // 启动流处理任务 spark.streams.awaitAnyTermination()
关键说明
- 该逻辑会自动处理所有层级的嵌套数组,比如
Shipment_Level2和Shipment_Level2_Level2RefNumbs都会生成独立的DataFrame - 每个子DataFrame会保留上层的关联字段(如
OrderId),方便后续关联查询 - 展平后的列名用下划线替换路径中的点,避免SQL语法冲突
内容的提问来源于stack exchange,提问作者CodeMonkey
相关产品推荐
相关产品推荐

