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

如何编写Spark SQL窗口函数UDAF计算分组排序后最后两行是否含true

解决方案

方案一:使用Spark内置函数实现(推荐)

无需自定义UDAF,直接通过窗口函数+聚合即可实现需求,性能更优且逻辑清晰:

SELECT 
  id,
  MAX(CASE WHEN isChurned = true THEN true ELSE false END) AS result
FROM (
  SELECT 
    id,
    isChurned,
    ROW_NUMBER() OVER (PARTITION BY id ORDER BY timestamp DESC) AS rn
  FROM your_table
) t
WHERE rn <= 2
GROUP BY id;

运行上述代码直接可以得到你需要的输出结果。


方案二:自定义UDAF实现

你遇到的merge方法顺序问题核心解决思路是:缓冲区不需要存储全量行,只需要存储当前最大的两个时间戳对应的isChurned值,合并时永远取所有条目中时间戳最大的两个即可,完全不需要关心行输入/合并的顺序。

完整实现代码(Scala)

import org.apache.spark.sql.expressions.UserDefinedAggregateFunction
import org.apache.spark.sql.types._
import org.apache.spark.sql.Row

class LastTwoHasTrueUDAF extends UserDefinedAggregateFunction {
  // 定义输入字段:时间戳、isChurned
  override def inputSchema: StructType = StructType(Seq(
    StructField("ts", LongType),
    StructField("is_churned", BooleanType)
  ))

  // 定义缓冲区结构:存储时间戳最大的两个条目的时间戳和对应isChurned值
  override def bufferSchema: StructType = StructType(Seq(
    StructField("ts1", LongType), // 最大时间戳
    StructField("churn1", BooleanType), // 最大时间戳对应isChurned
    StructField("ts2", LongType), // 第二大时间戳
    StructField("churn2", BooleanType) // 第二大时间戳对应isChurned
  ))

  override def dataType: DataType = BooleanType

  override def deterministic: Boolean = true

  // 初始化缓冲区
  override def initialize(buffer: org.apache.spark.sql.expressions.MutableAggregationBuffer): Unit = {
    buffer(0) = Long.MinValue
    buffer(1) = false
    buffer(2) = Long.MinValue
    buffer(3) = false
  }

  // 处理单条输入行
  override def update(buffer: org.apache.spark.sql.expressions.MutableAggregationBuffer, input: Row): Unit = {
    val currentTs = input.getAs[Long](0)
    // null值按false处理,可根据业务需求调整
    val currentChurn = Option(input.getAs[Boolean](1)).getOrElse(false)
    
    // 把当前行和缓冲区已有值合并,取时间戳最大的两个更新缓冲区
    val allEntries = Seq(
      (buffer.getAs[Long](0), buffer.getAs[Boolean](1)),
      (buffer.getAs[Long](2), buffer.getAs[Boolean](3)),
      (currentTs, currentChurn)
    )
    val sortedEntries = allEntries.sortBy(-_._1)
    buffer(0) = sortedEntries(0)._1
    buffer(1) = sortedEntries(0)._2
    buffer(2) = sortedEntries(1)._1
    buffer(3) = sortedEntries(1)._2
  }

  // 合并两个分区的缓冲区
  override def merge(buffer1: org.apache.spark.sql.expressions.MutableAggregationBuffer, buffer2: Row): Unit = {
    // 把两个缓冲区的4条记录全部取出,取时间戳最大的两个更新当前缓冲区
    val allEntries = Seq(
      (buffer1.getAs[Long](0), buffer1.getAs[Boolean](1)),
      (buffer1.getAs[Long](2), buffer1.getAs[Boolean](3)),
      (buffer2.getAs[Long](0), buffer2.getAs[Boolean](1)),
      (buffer2.getAs[Long](2), buffer2.getAs[Boolean](3))
    )
    val sortedEntries = allEntries.sortBy(-_._1)
    buffer1(0) = sortedEntries(0)._1
    buffer1(1) = sortedEntries(0)._2
    buffer1(2) = sortedEntries(1)._1
    buffer1(3) = sortedEntries(1)._2
  }

  // 输出最终结果
  override def evaluate(buffer: Row): Any = {
    buffer.getAs[Boolean](1) || buffer.getAs[Boolean](3)
  }
}

使用方式

首先注册UDAF:

spark.udf.register("last_two_has_true", new LastTwoHasTrueUDAF())

再执行SQL查询:

SELECT 
  id,
  last_two_has_true(timestamp, isChurned) AS result
FROM your_table
GROUP BY id;

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 20:36:04