如何编写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
相关产品推荐
相关产品推荐

