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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 06:07:19