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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 12:31:06