Scala通用Semigroup特质实现:多Case类聚合需求问询
为带聚合字段的Case类实现通用Semigroup特质
需求背景
我们有多个表示维度信息的Case类,每个类包含两类字段:
- 维度字段:无需聚合,只需保留其中一个实例的值(比如
Precipitation的date、GroceryList的groceryListId/groceryStoreName) - 统计字段:需要独立进行数值聚合(比如累加)
目前每个类都手动实现了combine方法来构建Semigroup,但希望定义一个通用的Semigroup特质,让任意这类Case类都能直接复用,避免重复代码。
现有手动实现代码
case class Precipitation(date: String, rainInches: Double, snowInches: Double) def combine(p1: Precipitation, p2: Precipitation): Precipitation = { p1.copy( rainInches = p1.rainInches + p2.rainInches, snowInches = p1.snowInches + p2.snowInches ) } case class GroceryList( groceryListId: String, groceryStoreName: String, eggs: Int, peaches: Int, flourPounds: Double ) def combine(g1: GroceryList, g2: GroceryList): GroceryList = { g1.copy( eggs = g1.eggs + g2.eggs, peaches = g1.peaches + g2.peaches, flourPounds = g1.flourPounds + g2.flourPounds ) }
通用Semigroup实现方案
方案1:Scala 3 利用Mirror元编程(推荐)
Scala 3的内省能力可以直接解析Case类的字段信息,我们通过标记特质区分聚合字段,实现通用的combine逻辑:
步骤1:定义聚合标记特质
// 用于标记需要聚合的字段类型 trait Aggregatable // 封装带标记的常用数值类型 type AggInt = Int with Aggregatable type AggDouble = Double with Aggregatable
步骤2:修改Case类标记聚合字段
case class Precipitation(date: String, rainInches: AggDouble, snowInches: AggDouble) case class GroceryList( groceryListId: String, groceryStoreName: String, eggs: AggInt, peaches: AggInt, flourPounds: AggDouble )
步骤3:实现通用Semigroup特质
trait Semigroup[A]: def combine(a1: A, a2: A): A object Semigroup: // 为所有Case类自动生成Semigroup实例 inline given [A <: Product](using mirror: Mirror.ProductOf[A]): Semigroup[A] = new Semigroup[A]: def combine(a1: A, a2: A): A = val fields1 = a1.productIterator.toList val fields2 = a2.productIterator.toList val combinedFields = fields1.zip(fields2).map { // 对标记为Aggregatable的字段执行累加 case (f1: Aggregatable, f2: Aggregatable) => (f1, f2) match case (i1: Int, i2: Int) => i1 + i2 case (d1: Double, d2: Double) => d1 + d2 case _ => f1 // 兜底逻辑,实际不会触发 // 非聚合字段保留第一个实例的值 case (f1, _) => f1 } // 从组合后的字段重建Case类实例 mirror.fromProduct(Tuple.fromArray(combinedFields.toArray))
使用方式
直接导入自动生成的Semigroup实例即可:
import Semigroup.given val p1 = Precipitation("2024-01-01", 1.2, 3.5) val p2 = Precipitation("2024-01-01", 0.8, 1.5) val combinedP = summon[Semigroup[Precipitation]].combine(p1, p2) // 结果:Precipitation(2024-01-01, 2.0, 5.0) val g1 = GroceryList("list1", "StoreA", 2, 5, 1.0) val g2 = GroceryList("list1", "StoreA", 3, 2, 0.5) val combinedG = summon[Semigroup[GroceryList]].combine(g1, g2) // 结果:GroceryList(list1,StoreA,5,7,1.5)
方案2:Scala 2 + Shapeless
如果使用Scala 2,可以借助Shapeless的泛型能力实现通用逻辑:
步骤1:添加Shapeless依赖
libraryDependencies += "com.chuusai" %% "shapeless" % "2.3.10"
步骤2:实现通用Semigroup
import shapeless._ import shapeless.labelled.FieldType import shapeless.ops.hlist.{Mapper, ZipWith} // 聚合标记特质 trait Aggregatable // 定义字段聚合逻辑的Poly object FieldAggregator extends Poly2 { // 对标记为Aggregatable的数值字段执行累加 implicit def aggregatableInt[K]: Case.Aux[FieldType[K, Int with Aggregatable], FieldType[K, Int with Aggregatable], FieldType[K, Int with Aggregatable]] = at((a, b) => field[K](a + b)) implicit def aggregatableDouble[K]: Case.Aux[FieldType[K, Double with Aggregatable], FieldType[K, Double with Aggregatable], FieldType[K, Double with Aggregatable]] = at((a, b) => field[K](a + b)) // 非聚合字段保留第一个实例的值 implicit def nonAggregatable[K, V]: Case.Aux[FieldType[K, V], FieldType[K, V], FieldType[K, V]] = at((a, _) => a) } trait Semigroup[A] { def combine(a1: A, a2: A): A } object Semigroup { def apply[A](implicit sg: Semigroup[A]): Semigroup[A] = sg // 为Case类自动生成Semigroup实例 implicit def genericSemigroup[A, L <: HList]( implicit gen: LabelledGeneric.Aux[A, L], zipWith: ZipWith.Aux[L, L, FieldAggregator.type, L] ): Semigroup[A] = new Semigroup[A] { def combine(a1: A, a2: A): A = gen.from(zipWith(gen.to(a1), gen.to(a2))) } }
使用方式
给Case类的统计字段标记聚合特质后即可使用:
case class Precipitation(date: String, rainInches: Double with Aggregatable, snowInches: Double with Aggregatable) case class GroceryList( groceryListId: String, groceryStoreName: String, eggs: Int with Aggregatable, peaches: Int with Aggregatable, flourPounds: Double with Aggregatable ) val p1 = Precipitation("2024-01-01", 1.2, 3.5) val p2 = Precipitation("2024-01-01", 0.8, 1.5) val combinedP = Semigroup[Precipitation].combine(p1, p2)
扩展说明
- 可以在
combine方法中添加维度字段一致性校验,避免不同维度的数据被错误聚合 - 若需要支持其他聚合逻辑(比如取最大值、乘法),只需修改聚合字段的处理逻辑即可
内容的提问来源于stack exchange,提问作者Michael K
相关产品推荐
相关产品推荐

