Scala多阶段元编程中如何匹配Array[Int]子表达式?
解决Scala多阶段元编程中数组点积的编译期匹配问题
问题根源
你实现的ArrayFromExpr中,非空数组的模式匹配逻辑不符合Scala宏AST的结构,导致无法从Expr[Array[Int]]中提取编译期常量数组,最终v1.value和v2.value返回None,进入case (None, None)分支返回99,而非预期的异常。
修复步骤
1. 修正FromExpr[Array[Int]]的实现
原代码的非空数组匹配模式错误,正确的做法是直接匹配Varargs包裹的元素表达式序列,逐个提取每个元素的编译期Int值:
import scala.quoted.* given ArrayFromExpr: FromExpr[Array[Int]] with def unapply(x: Expr[Array[Int]])(using Quotes): Option[Array[Int]] = x match // 匹配空数组 case '{ Array[Int]() } => Some(Array.empty[Int]) // 匹配带元素的数组:提取每个元素的编译期值 case '{ Array[Int](${Varargs(exprs)}: _*) } => val elements = exprs.map { case Expr(value: Int) => value // 若元素无法提取编译期值,直接返回None case _ => return None } Some(elements.toArray) // 其他情况返回None case _ => None
2. 调整宏实现的错误处理逻辑
原代码在宏中抛出RuntimeException会导致编译失败(宏代码在编译期执行),但你的测试是运行时断言,逻辑矛盾。需区分编译期常量和非常量数组的处理:
def dotImpl(v1: Expr[Array[Int]], v2: Expr[Array[Int]])(using q: Quotes): Expr[Int] = { import q.reflect.* (v1.value, v2.value) match { case (Some(arr1), Some(arr2)) if arr1.length != arr2.length => // 编译期报错并终止编译 report.errorAndAbort("Cannot compute dot product of arrays having different lengths.") case (Some(arr1), Some(arr2)) if arr1.isEmpty => '{ 0 } case (Some(arr1), Some(arr2)) => // 编译期计算点积,返回对应的Expr val dot = arr1.zip(arr2).map(_ * _).sum Expr(dot) case _ => // 非编译期常量数组,生成运行时检查代码 '{ val a1 = $v1 val a2 = $v2 if (a1.length != a2.length) throw new RuntimeException("Cannot compute dot product of arrays having different lengths.") else a1.zip(a2).map(_ * _).sum } } }
3. 修正测试逻辑
编译期检查的场景需验证编译失败,运行时检查的场景验证运行时异常:
"dot product" should { "throw compile error for different length constant arrays" in { val error = compileError(""" val v1 = Array[Int](2, 1) val v2 = Array[Int](1, 2, 3) dot(v1, v2) """).message error should include("Cannot compute dot product of arrays having different lengths.") } "throw runtime exception for different length non-constant arrays" in { def getV1 = Array[Int](2, 1) def getV2 = Array[Int](1, 2, 3) a [RuntimeException] should be thrownBy { dot(getV1, getV2) } } "return correct value for same length arrays" in { val v1 = Array[Int](2, 3) val v2 = Array[Int](4, 5) dot(v1, v2) shouldBe (2*4 + 3*5) // 23 } }
关键说明
Varargs(exprs)是匹配数组构造参数的标准方式,原代码的Exprs(elements)并非Quotes的合法模式语法。- 宏需区分编译期常量和非常量场景:常量数组直接在编译期计算结果,非常量数组生成运行时检查代码。
- 编译期报错使用
report.errorAndAbort,直接终止编译;运行时异常则在生成的代码中抛出。
内容的提问来源于stack exchange,提问作者Wuj
相关产品推荐
相关产品推荐

