Scala3中如何用宏从模式匹配提取所有样例类?
Scala3 提取模式匹配中样例类的简便实现
在Scala3中,无需手动匹配所有Tree类型,你可以使用TreeTraverser自动递归遍历语法树,只需关注目标节点(比如CaseDef)即可,思路和Scala2中Transformer类似,且适配了Scala3的API。
具体实现代码
基于Scala3的宏API实现extract方法:
import scala.quoted.* import scala.reflect.api.TreeTraverser sealed trait Adt case class A(i: Int) extends Adt case class B(i: Int) extends Adt inline def extract(pf: PartialFunction[Adt, Unit]): List[Class[_ <: Adt]] = ${ extractImpl('pf) } def extractImpl(pfExpr: Expr[PartialFunction[Adt, Unit]])(using Quotes): Expr[List[Class[_ <: Adt]]] = { import quotes.reflect.* val caseClassSymbols = collection.mutable.ListBuffer[Symbol]() // 自定义遍历器,自动递归处理所有子树 val traverser = new TreeTraverser { override def traverseTree(tree: Tree)(owner: Symbol): Unit = { tree match { // 只处理模式匹配的CaseDef节点 case CaseDef(pattern, _, _) => pattern match { // 匹配样例类的构造模式,提取类符号 case Apply(TypeApply(Select(New(TypeIdent(clsSym)), _), _), _) => caseClassSymbols += clsSym case _ => // 忽略非样例类的模式 } case _ => // 其他节点不做特殊处理,交给父类继续遍历子节点 } super.traverseTree(tree)(owner) } } // 开始遍历传入的PartialFunction语法树 traverser.traverseTree(pfExpr.asTerm)(Symbol.spliceOwner) // 将收集到的类符号转为Class实例的表达式 val classExprs = caseClassSymbols.map { sym => Expr(sym.companionModule.moduleClass.javaClass).asInstanceOf[Expr[Class[_ <: Adt]]] } Expr.ofList(classExprs) }
使用方式
和Scala2的调用方式完全一致:
val result = extract { case A(i) => println(i) case B(i) => println(i) } // result 为 List(class A, class B)
核心优势
TreeTraverser会自动处理所有子节点的递归遍历,无需手动匹配Inlined、Block等各种Tree类型,大幅简化代码。- 只需聚焦于
CaseDef和样例类模式的匹配逻辑,其余遍历工作由父类方法自动完成。
内容的提问来源于stack exchange,提问作者KrzyH
相关产品推荐
相关产品推荐

