测试Spark函数时,如何正常执行spark.table并mock指定的read操作?
解决Spark私有函数测试中Mock指定Read操作的问题
问题背景
需要测试如下私有函数:
private HashMap<String, Dataset<Row>> getDataSources(SparkSession spark) { HashMap<String, Dataset<Row>> ds = new HashMap<String, Dataset<Row>>(); Dataset<Row>dimTenant = spark.table(dbName + "." + SparkConstants.DIM_TENANT) .select("tenant_key", "tenant_id"); Map<String, String> options = new HashMap<>(); options.put("table", bookValueTable); options.put("zkUrl", zkUrl); Dataset<Row> bookValue = spark.read().format("org.apache.phoenix.spark") .options(options) .load(); ds.put("dimTenant", dimTenant); ds.put("bookValue", bookValue); return ds; }
测试需求:
- 保留
spark.table()的真实执行逻辑 - 仅针对
spark.read().format("org.apache.phoenix.spark").options(options).load()的调用Mock输出,不影响其他spark.read()相关操作(比如spark.table()内部调用的spark.read().table)
此前尝试的方案存在问题:
- 深度Mock
DataframeReader会连带影响spark.table()的内部逻辑 - Spy
spark.read()对象无效,因为每次调用spark.read()都会生成新实例
解决方案1:JDK动态代理实现精准拦截
通过自定义DataframeReader的代理类,仅拦截目标格式的链式调用,其余操作走真实逻辑。
步骤1:实现代理类
import org.apache.spark.sql.DataFrameReader; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; import java.lang.reflect.Proxy; import java.util.Map; public class MockableDataframeReaderProxy implements InvocationHandler { private final DataframeReader realReader; private final Dataset<Row> mockPhoenixDataset; private final Map<String, String> expectedOptions; public MockableDataframeReaderProxy(DataframeReader realReader, Dataset<Row> mockPhoenixDataset, Map<String, String> expectedOptions) { this.realReader = realReader; this.mockPhoenixDataset = mockPhoenixDataset; this.expectedOptions = expectedOptions; } @Override public Object invoke(Object proxy, Method method, Object[] args) throws Throwable { if ("format".equals(method.getName()) && args != null && args.length == 1) { String format = (String) args[0]; if ("org.apache.phoenix.spark".equals(format)) { return Proxy.newProxyInstance( DataframeReader.class.getClassLoader(), new Class[]{DataframeReader.class}, new FormatChainedInvocationHandler(realReader, mockPhoenixDataset, expectedOptions) ); } } return method.invoke(realReader, args); } private static class FormatChainedInvocationHandler implements InvocationHandler { private final DataframeReader realReader; private final Dataset<Row> mockPhoenixDataset; private final Map<String, String> expectedOptions; private Map<String, String> actualOptions; public FormatChainedInvocationHandler(DataframeReader realReader, Dataset<Row> mockPhoenixDataset, Map<String, String> expectedOptions) { this.realReader = realReader; this.mockPhoenixDataset = mockPhoenixDataset; this.expectedOptions = expectedOptions; } @Override public Object invoke(Object proxy, Method method, Object[] args) throws Throwable { if ("options".equals(method.getName()) && args != null && args.length == 1) { this.actualOptions = (Map<String, String>) args[0]; // 可选:验证传入的options是否符合预期 assert actualOptions.equals(expectedOptions); return proxy; } else if ("load".equals(method.getName())) { return mockPhoenixDataset; } return method.invoke(realReader, args); } } public static DataframeReader createProxy(DataframeReader realReader, Dataset<Row> mockPhoenixDataset, Map<String, String> expectedOptions) { return (DataframeReader) Proxy.newProxyInstance( DataframeReader.class.getClassLoader(), new Class[]{DataframeReader.class}, new MockableDataframeReaderProxy(realReader, mockPhoenixDataset, expectedOptions) ); } }
步骤2:在测试中使用代理
import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.SparkSession; import org.junit.Test; import org.mockito.Mockito; import java.util.HashMap; import java.util.Map; import static org.mockito.Mockito.when; public class YourClassTest { @Test public void testGetDataSources() throws Exception { // 初始化真实SparkSession,用于执行spark.table的真实逻辑 SparkSession realSpark = SparkSession.builder() .master("local[1]") .appName("Test") .getOrCreate(); // 构造Mock的Phoenix数据集 Dataset<Row> mockBookValue = realSpark.createDataFrame( // 填入测试数据和对应Schema // 示例:List.of(RowFactory.create("test", 123)), // StructType.fromDDL("tenant_key string, value int") ); // 定义预期的options参数 Map<String, String> expectedOptions = new HashMap<>(); expectedOptions.put("table", bookValueTable); expectedOptions.put("zkUrl", zkUrl); // Spy真实SparkSession,保留table方法的原生逻辑 SparkSession sparkSpy = Mockito.spy(realSpark); // 创建代理的DataframeReader DataframeReader proxyReader = MockableDataframeReaderProxy.createProxy( realSpark.read(), mockBookValue, expectedOptions ); // 让Spy的read方法返回代理实例 when(sparkSpy.read()).thenReturn(proxyReader); // 反射调用私有函数getDataSources YourClass target = new YourClass(); // 确保target的dbName等成员变量已初始化 HashMap<String, Dataset<Row>> result = (HashMap<String, Dataset<Row>>) YourClass.class.getDeclaredMethod("getDataSources", SparkSession.class) .invoke(target, sparkSpy); // 验证结果:dimTenant为真实查询结果,bookValue为Mock数据集 // 示例:assert result.get("dimTenant").count() > 0; // assert result.get("bookValue").equals(mockBookValue); } }
解决方案2:PowerMock拦截Read实例创建
如果依赖PowerMock,可直接拦截SparkSession.read()的调用,返回同一个Spy的DataframeReader,并Mock指定链式调用。
import org.apache.spark.sql.DataFrameReader; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.SparkSession; import org.junit.Test; import org.junit.runner.RunWith; import org.mockito.Mockito; import org.powermock.api.mockito.PowerMockito; import org.powermock.core.classloader.annotations.PrepareForTest; import org.powermock.modules.junit4.PowerMockRunner; @RunWith(PowerMockRunner.class) @PrepareForTest(SparkSession.class) public class YourClassPowerMockTest { @Test public void testGetDataSources() throws Exception { SparkSession realSpark = SparkSession.builder() .master("local[1]") .appName("Test") .getOrCreate(); Dataset<Row> mockBookValue = realSpark.createDataFrame(...); // Spy真实的DataframeReader DataframeReader readerSpy = Mockito.spy(realSpark.read()); // Mock目标格式的链式调用 DataframeReader mockFormatReader = Mockito.mock(DataframeReader.class); Mockito.when(readerSpy.format("org.apache.phoenix.spark")).thenReturn(mockFormatReader); Mockito.when(mockFormatReader.options(Mockito.anyMap())).thenReturn(mockFormatReader); Mockito.when(mockFormatReader.load()).thenReturn(mockBookValue); // 拦截SparkSession.read(),每次返回同一个Spy实例 PowerMockito.when(realSpark.read()).thenReturn(readerSpy); // 反射调用私有方法并验证 YourClass target = new YourClass(); HashMap<String, Dataset<Row>> result = (HashMap<String, Dataset<Row>>) YourClass.class.getDeclaredMethod("getDataSources", SparkSession.class) .invoke(target, realSpark); // 验证逻辑... } }
方案对比
- 动态代理方案:无需额外依赖,仅用JDK反射+Mockito,兼容性好,适合对依赖有严格限制的场景
- PowerMock方案:代码更简洁,但需要引入PowerMock依赖,可能与部分测试框架存在兼容性问题
内容的提问来源于stack exchange,提问作者ujjawal
相关产品推荐
相关产品推荐

