Spark中实现非终止式CASE WHEN:累加匹配结果为数组
解决方法
CASE WHEN是短路求值逻辑,只会返回第一个满足条件的结果,要收集所有匹配项,得把每个条件单独判断生成对应结果,再过滤无效值后整理成目标格式。
方法1:直接用Spark内置函数实现
通过array()打包所有条件的结果(不满足则为null),再用array_filter()剔除null元素,最后根据结果数量决定返回数组还是单个值:
from pyspark.sql import functions as F df_out = df_source.withColumn( "CalculatedValue", F.expr(""" let matches = array_filter( array( CASE WHEN Col1 = 100 THEN 'AAA111' END, CASE WHEN Col2 = 'aaa' THEN 'BBB222' END, CASE WHEN Col3 = 'zzz' THEN 'CCC333' END ), x -> x IS NOT NULL ) in if size(matches) = 1 then matches[0] else matches """) )
方法2:适配动态构建场景
如果你的条件是动态生成的,可以先定义条件列表,再拼接成表达式字符串:
from pyspark.sql import functions as F # 定义动态条件(可根据实际需求扩展) conditions = [ ("Col1 = 100", "'AAA111'"), ("Col2 = 'aaa'", "'BBB222'"), ("Col3 = 'zzz'", "'CCC333'") ] # 生成单个条件的CASE片段 case_clauses = [f"CASE WHEN {cond} THEN {val} END" for cond, val in conditions] # 拼接成完整表达式 expr_str = f""" let matches = array_filter(array({','.join(case_clauses)}), x -> x IS NOT NULL) in if size(matches) = 1 then matches[0] else matches """ df_out = df_source.withColumn("CalculatedValue", F.expr(expr_str))
最终输出
执行后会得到符合需求的结果:
| Id | Col1 | Col2 | Col3 | CalculatedValue |
|---|---|---|---|---|
| 1 | 100 | aaa | xxx | [AAA111, BBB222] |
| 2 | 200 | aaa | yyy | BBB222 |
| 3 | 300 | ccc | zzz | CCC333 |
补充说明
- 如果不需要单个结果显示为字符串,直接返回
matches即可,此时单个匹配项会以长度为1的数组形式输出 - 动态构建逻辑可以灵活扩展,只需往
conditions列表中添加新的条件和对应结果即可
内容的提问来源于stack exchange,提问作者Martin
相关产品推荐
相关产品推荐

