Spark Scala:基于窗口函数实现节假日匹配后的日期调整需求
Spark Scala 实现节假日日期替换(广播变量+UDF方案)
需求说明
给定两个DataFrame:
holiday_df:存储不同国家/货币的节假日日期everyday_df:存储日常日期记录
需要按国家代码和货币代码匹配,将everyday_df中属于节假日的date_demo替换为下一个非节假日工作日(该日期不能出现在holiday_df中)。
示例数据
holiday_df
| Country_code | currency_code | date |
|---|---|---|
| Gb | gbp | 2022-04-15 |
| Gb | gbp | 2022-04-16 |
| US | usd | 2022-04-17 |
| Gb | gbp | 2022-04-18 |
| Gb | gbp | 2022-04-21 |
everyday_df
| Country_code_demo | currency_code_demo | date_demo |
|---|---|---|
| Gb | gbp | 2022-04-14 |
| Gb | gbp | 2022-04-15 |
| Gb | gbp | 2022-04-16 |
| Gb | gbp | 2022-04-18 |
实现方案
采用广播变量存储节假日集合(避免重复查询)+ 自定义UDF计算下一个非节假日的方案,结合Spark分布式特性保证执行效率。
完整代码
import org.apache.spark.sql.SparkSession import org.apache.spark.sql.functions._ object HolidayDateAdjuster { def main(args: Array[String]): Unit = { val spark = SparkSession.builder() .appName("HolidayDateAdjustment") .master("local[*]") .getOrCreate() import spark.implicits._ // 构建节假日DataFrame val holidayData = Seq( ("Gb", "gbp", "2022-04-15"), ("Gb", "gbp", "2022-04-16"), ("US", "usd", "2022-04-17"), ("Gb", "gbp", "2022-04-18"), ("Gb", "gbp", "2022-04-21") ).map(row => (row._1, row._2, java.sql.Date.valueOf(row._3))) .toDF("Country_code", "currency_code", "date") // 构建日常日期DataFrame val everydayData = Seq( ("Gb", "gbp", "2022-04-14"), ("Gb", "gbp", "2022-04-15"), ("Gb", "gbp", "2022-04-16"), ("Gb", "gbp", "2022-04-18") ).map(row => (row._1, row._2, java.sql.Date.valueOf(row._3))) .toDF("Country_code_demo", "currency_code_demo", "date_demo") // 生成并广播每个国家+货币对应的节假日集合 val holidayMap = holidayData.groupBy("Country_code", "currency_code") .agg(collect_set("date").as("holidays")) .rdd.map { row => val key = (row.getAs[String]("Country_code"), row.getAs[String]("currency_code")) val holidays = row.getAs[Seq[java.sql.Date]]("holidays").toSet (key, holidays) }.collectAsMap() val broadcastHolidays = spark.sparkContext.broadcast(holidayMap) // 定义UDF:计算下一个非节假日工作日 val getNextNonHoliday = udf((country: String, currency: String, currentDate: java.sql.Date) => { val holidays = broadcastHolidays.value.getOrElse((country, currency), Set.empty[java.sql.Date]) if (!holidays.contains(currentDate)) { currentDate } else { var nextDay = new java.sql.Date(currentDate.getTime + 86400000L) // 增加一天(毫秒数) while (holidays.contains(nextDay)) { nextDay = new java.sql.Date(nextDay.getTime + 86400000L) } nextDay } }) // 计算调整后的日期 val resultDF = everydayData.withColumn( "date_updated", getNextNonHoliday($"Country_code_demo", $"currency_code_demo", $"date_demo") ) // 展示结果 resultDF.show(false) } }
代码说明
- 数据初始化:将示例数据转换为Spark DataFrame,强制日期类型为
DateType,避免字符串解析错误。 - 广播节假日集合:将每个国家/货币对应的节假日日期收集为Set,通过广播变量分发到所有节点,避免重复查询,大幅提升分布式场景下的性能。
- 自定义UDF:核心逻辑判断当前日期是否为节假日,若是则循环加一天,直到找到第一个不在节假日集合中的日期。
- 生成结果:调用UDF生成
date_updated列,得到最终调整后的日期。
预期输出
+-----------------+-------------------+----------+------------+ |Country_code_demo|currency_code_demo |date_demo |date_updated| +-----------------+-------------------+----------+------------+ |Gb |gbp |2022-04-14|2022-04-14 | |Gb |gbp |2022-04-15|2022-04-17 | |Gb |gbp |2022-04-16|2022-04-17 | |Gb |gbp |2022-04-18|2022-04-19 | +-----------------+-------------------+----------+------------+
内容的提问来源于stack exchange,提问作者Vaibhav Kulkarni
相关产品推荐
相关产品推荐

