You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于Scala实现返回Array类型的Spark UDAF:按时间排序渠道

实现按Time升序排序Channel的Spark UDAF(Scala)

没问题,我帮你搞定这个需求——我们需要写一个自定义聚合函数(UDAF),按id分组后,把channel字段按照time的升序排列,最终返回Array[String]类型,而且支持Spark SQL直接调用。

步骤1:编写UDAF类

我们继承UserDefinedAggregateFunction(这个类在SQL场景下适配性更好),核心逻辑是在聚合过程中收集每个分组的channel和time对,最后排序并提取channel组成数组:

import org.apache.spark.sql.expressions.UserDefinedAggregateFunction
import org.apache.spark.sql.types._
import org.apache.spark.sql.{Row, SparkSession}

class SortChannelsByTimeUDAF extends UserDefinedAggregateFunction {
  // 输入数据的结构:(channel字符串, time整数)
  override def inputSchema: StructType = StructType(
    StructField("channel", StringType) :: StructField("time", IntegerType) :: Nil
  )

  // 缓冲区的结构:存储收集到的(channel, time)元组数组
  override def bufferSchema: StructType = StructType(
    StructField("pairs", ArrayType(StructType(
      StructField("channel", StringType) :: StructField("time", IntegerType) :: Nil
    ))) :: Nil
  )

  // 返回值类型:字符串数组
  override def dataType: DataType = ArrayType(StringType)

  // 确定性标记:相同输入必然返回相同输出
  override def deterministic: Boolean = true

  // 初始化缓冲区:空数组
  override def initialize(buffer: Row): Unit = {
    buffer.update(0, Array.empty[(String, Int)])
  }

  // 更新缓冲区:把当前输入的(channel, time)添加到缓冲区数组
  override def update(buffer: Row, input: Row): Unit = {
    val currentPairs = buffer.getAs[Array[(String, Int)]](0)
    val newPair = (input.getAs[String](0), input.getAs[Int](1))
    buffer.update(0, currentPairs :+ newPair)
  }

  // 合并缓冲区:把两个分区的缓冲区数组合并
  override def merge(buffer1: Row, buffer2: Row): Unit = {
    val pairs1 = buffer1.getAs[Array[(String, Int)]](0)
    val pairs2 = buffer2.getAs[Array[(String, Int)]](0)
    buffer1.update(0, pairs1 ++ pairs2)
  }

  // 计算最终结果:按time升序排序,提取channel组成数组
  override def evaluate(buffer: Row): Any = {
    val pairs = buffer.getAs[Array[(String, Int)]](0)
    pairs.sortBy(_._2).map(_._1)
  }
}

步骤2:注册UDAF并执行SQL查询

接下来把UDAF注册到SparkSession,然后用SQL完成分组排序:

// 初始化SparkSession(本地测试用,生产环境去掉master配置)
val spark = SparkSession.builder()
  .appName("SortChannelsUDAFTest")
  .master("local[*]")
  .getOrCreate()

// 创建你的测试DataFrame
val myDF = Seq(
  (1,"A",100), (1,"E",300), (1,"B",200),
  (2,"A",200), (2,"C",300), (2,"D",100)
).toDF("id","channel","time")

// 注册自定义聚合函数
spark.udf.register("sort_channels_by_time", new SortChannelsByTimeUDAF())

// 执行Spark SQL查询
spark.sql("""
  SELECT id, sort_channels_by_time(channel, time) AS sorted_channels
  FROM myDF
  GROUP BY id
""").show(false)

预期输出

运行后会得到如下结果:

+---+----------------+
|id |sorted_channels |
+---+----------------+
|1  |[A, B, E]       |
|2  |[D, A, C]       |
+---+----------------+

补充说明

  • 如果你的Spark版本是3.0+,也可以用Aggregator实现类型更安全的UDAF,但上面的UserDefinedAggregateFunction在SQL中调用更直接。
  • 缓冲区用Array而非ListBuffer是因为Spark的Row只能存储序列化后的集合类型,Array更适配Spark内部的序列化机制。

内容的提问来源于stack exchange,提问作者Regressor

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.22 07:52:21