如何用ScalaTest与ScalaMock实现类方法的行为单元测试?
ScalaTest + ScalaMock 测试实践方案
问题1:验证公共方法c对私有方法的调用逻辑
场景
类A包含私有方法a、b,以及公共方法c。c会根据传入参数,通过if分支调用a或b。a、b已完成单元测试,需验证c的分支调用逻辑是否正确。
实现方案
使用ScalaMock的**间谍(Spy)**包装真实实例,跟踪私有方法的调用情况,验证c在特定参数下是否触发了预期的私有方法。
代码示例
import org.scalatest.funsuite.AnyFunSuite import org.scalamock.scalatest.MockFactory // 待测试类A class A { private def a(): Unit = {} private def b(): Unit = {} def c(param: String): Unit = { if (param == "trigger-a") a() else b() } } // 测试类 class ClassATest extends AnyFunSuite with MockFactory { test("方法c传入trigger-a时调用私有方法a") { val aInstance = spy(new A()) aInstance.c("trigger-a") // 验证a被调用1次,b未被调用 (aInstance.a _).verify().once() (aInstance.b _).verify().never() } test("方法c传入非trigger-a参数时调用私有方法b") { val aInstance = spy(new A()) aInstance.c("trigger-b") // 验证b被调用1次,a未被调用 (aInstance.b _).verify().once() (aInstance.a _).verify().never() } }
说明
spy会保留类A的真实逻辑,同时跟踪方法调用记录。- ScalaMock通过反射机制支持私有方法的调用验证,无需修改原类的访问权限。
问题2:测试受保护的final方法generate
场景
目标类包含受保护的final方法generate,其依赖的nextOrExpand、usePoint等私有方法已完成测试,需覆盖generate的所有分支逻辑。
实现方案
- 创建测试子类,将
generate包装为公共方法(若测试类与目标类不同包); - 针对
generate的每个分支(空/非空分支、排序比较结果等)编写测试用例; - 用Spy跟踪依赖方法的调用,验证分支逻辑的正确性,或直接验证异常抛出行为。
代码示例
import org.scalatest.funsuite.AnyFunSuite import org.scalamock.scalatest.MockFactory // 自定义异常 case class SamePointException() extends Exception case class LowerRestrictionException() extends Exception // 待测试的目标类 class TreeGenerator[L: Ordering, C] { protected val leafOrdering: Ordering[L] = implicitly[Ordering[L]] // 已测试的私有依赖方法 private def nextOrExpand(p: L)(f: L => C): C = f(p) private def usePoint(p: L): C = s"usePoint($p)" private def expand(p: L): C = s"expand($p)" private def minimalRestricted(left: List[L], start: L => List[L], right: List[L], rp: L): C = "minimalRestricted" private def next(left: List[L], lp: L): C = "next" // 待测试的protected final方法 protected final def generate(leftBranch: List[L], leftPoint: L, rightBranch: List[L], rightPoint: L): C = { val lEmpty = leftBranch.isEmpty val rEmpty = rightBranch.isEmpty if (lEmpty && rEmpty) { val cc = leafOrdering.compare(leftPoint, rightPoint) if (cc < 0) { nextOrExpand(leftPoint){ p => if (leafOrdering.compare(p, rightPoint) < 0) usePoint(p) else expand(leftPoint) } } else if (cc == 0) throw SamePointException() else throw LowerRestrictionException() } else if (lEmpty) { val cc = leafOrdering.compare(leftPoint, rightBranch.head) if (cc < 0) nextOrExpand(leftPoint){ p => if (leafOrdering.compare(p, rightBranch.head) <= 0) usePoint(p) else expand(leftPoint) } else if (cc == 0) minimalRestricted(List(leftPoint), List(_), rightBranch.tail, rightPoint) else throw LowerRestrictionException() } else if (rEmpty) if (leafOrdering.compare(leftBranch.head, rightPoint) < 0) next(leftBranch, leftPoint) else throw LowerRestrictionException() else { if (leafOrdering.compare(leftBranch.head, rightBranch.head) < 0) next(leftBranch, leftPoint) else throw LowerRestrictionException() } } } // 测试类 class TreeGeneratorTest extends AnyFunSuite with MockFactory { // 测试子类:将protected方法暴露为public class TestableTreeGenerator[L: Ordering, C] extends TreeGenerator[L, C] { def publicGenerate(leftBranch: List[L], leftPoint: L, rightBranch: List[L], rightPoint: L): C = { generate(leftBranch, leftPoint, rightBranch, rightPoint) } } test("左右分支均为空且leftPoint < rightPoint时,执行nextOrExpand并触发usePoint/expand") { val generator = spy(new TestableTreeGenerator[Int, String]) val leftPoint = 1 val rightPoint = 2 generator.publicGenerate(Nil, leftPoint, Nil, rightPoint) // 验证nextOrExpand被调用,且传入的函数逻辑符合预期 (generator.nextOrExpand _).verify(leftPoint).once().onCall { (p, f) => assert(f(1) == "usePoint(1)") assert(f(2) == "expand(1)") "test-result" } } test("左右分支均为空且leftPoint == rightPoint时,抛出SamePointException") { val generator = new TestableTreeGenerator[Int, String] val point = 1 assertThrows[SamePointException] { generator.publicGenerate(Nil, point, Nil, point) } } test("左分支为空且leftPoint > rightBranch.head时,抛出LowerRestrictionException") { val generator = new TestableTreeGenerator[Int, String] val leftPoint = 3 val rightBranch = List(2) val rightPoint = 4 assertThrows[LowerRestrictionException] { generator.publicGenerate(Nil, leftPoint, rightBranch, rightPoint) } } test("右分支为空且leftBranch.head < rightPoint时,调用next方法") { val generator = spy(new TestableTreeGenerator[Int, String]) val leftBranch = List(1) val leftPoint = 0 val rightPoint = 2 generator.publicGenerate(leftBranch, leftPoint, Nil, rightPoint) (generator.next _).verify(leftBranch, leftPoint).once() } }
说明
- 针对
protected方法,通过测试子类暴露为公共方法,避免修改原类的访问控制; - 对于异常分支,直接调用方法并验证是否抛出预期异常;
- 对于依赖方法的调用验证,使用Spy跟踪调用记录,确保分支逻辑执行了正确的依赖方法。
内容的提问来源于stack exchange,提问作者andrey.ladniy
相关产品推荐
相关产品推荐

