Mockito thenReturn在RichSourceFunction中失效问题排查与解决
问题分析
你的问题核心在于Flink的RichSourceFunction会被序列化传输,而默认的Mockito Mock对象不支持序列化。当你把Source传入Flink环境后,Flink会对Source进行序列化,再在执行线程中反序列化,此时拿到的Fetcher实例已经不是你最初配置Stub的那个Mock对象了。这就导致run()方法中调用getItems()时,会执行原方法逻辑,触发未初始化字段的空指针异常;而构造函数中Mock生效,是因为当时还没进入Flink的序列化流程,用的是原始Mock实例。
解决方案
1. 创建可序列化的Mock对象
修改Mock对象的创建方式,指定序列化配置,让Mock能在序列化/反序列化后保留Stub规则:
Fetcher mockFetcher = Mockito.mock(Fetcher.class, Mockito.withSettings().serializable());
这样run()方法中调用fetcher.getItems()时,会直接返回你设置的预设结果,不会执行原方法逻辑。
2. 缩短测试休眠时间(建议)
原代码中休眠5分钟(300000ms)会导致测试等待过久,建议给Source增加自定义间隔的构造参数,测试时传入短间隔:
修改Source类:
public static class Source extends RichSourceFunction<Set<Integer>> { private boolean isRunning; private final Fetcher fetcher; private final long sleepInterval; // 新增带间隔的构造方法 public Source(final Fetcher fetcher, long sleepInterval) { this.fetcher = fetcher; this.isRunning = true; this.sleepInterval = sleepInterval; } // 保留原构造方法兼容旧逻辑 public Source(final Fetcher fetcher) { this(fetcher, 300000); } @Override public void run(SourceContext<Set<Integer>> ctx) throws Exception { while (this.isRunning) { try { Set<Integer> items = fetcher.getItems(); if (items != null) { ctx.collect(items); } TimeUnit.MILLISECONDS.sleep(sleepInterval); } catch (Exception e) { System.out.println("exception"); } } } @Override public void cancel() { this.isRunning = false; } }
测试时传入短间隔:
Head.Source source = new Head.Source(mockFetcher, 100); // 休眠100ms,快速获取两次结果
3. 增加测试超时保护(可选)
避免因意外情况导致测试挂起,给测试方法加上超时注解:
@Test(timeout = 5000) // 5秒超时 public void testSource() throws Exception { // ... 测试逻辑 }
修改后的完整测试代码
@Test(timeout = 5000) public void testSource() throws Exception { // 创建可序列化的Mock对象 Fetcher mockFetcher = Mockito.mock(Fetcher.class, Mockito.withSettings().serializable()); Set<Integer> items1 = new HashSet<>(Arrays.asList(1, 3, 5)); Set<Integer> items2 = new HashSet<>(Collections.singletonList(7)); Mockito.when(mockFetcher.getItems()) .thenReturn(items1) .thenReturn(items2); // 传入短休眠间隔的Source Head.Source source = new Head.Source(mockFetcher, 100); DataStream<Set<Integer>> items = env.addSource(source); Iterator<Set<Integer>> iterator = items.executeAndCollect(); Iterable<Set<Integer>> iterable = () -> iterator; List<Set<Integer>> collectedItems = StreamSupport.stream(iterable.spliterator(), false).limit(2).collect(Collectors.toList()); Assert.assertEquals(2, collectedItems.size()); // 验证Mock方法被调用两次 Mockito.verify(mockFetcher, Mockito.times(2)).getItems(); }
内容的提问来源于stack exchange,提问作者Scott
相关产品推荐
相关产品推荐

