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

如何用MockConsumer单元测试Spring Boot的MessageListenerContainer

如何用Mock Kafka Consumer测试Spring Kafka的MessageListenerContainer

核心思路是替换真实的ConsumerFactory为基于Apache Kafka MockConsumer的实现,手动构建容器并模拟消息,无需启动完整Spring上下文。

实现步骤

  1. 自定义ConsumerFactory,返回MockConsumer实例而非真实消费者
  2. 手动构建ConcurrentMessageListenerContainer,注入自定义ConsumerFactory和你的消息监听器
  3. 给MockConsumer添加模拟消息记录,启动容器触发消费逻辑
  4. 验证消息监听器是否正确处理所有消息

Java/JUnit 示例

假设你的消息监听器实现了MessageListener<String, String>:

import org.apache.kafka.clients.consumer.ConsumerRecord;
import org.apache.kafka.clients.consumer.MockConsumer;
import org.apache.kafka.clients.consumer.OffsetResetStrategy;
import org.apache.kafka.common.TopicPartition;
import org.junit.jupiter.api.Test;
import org.springframework.kafka.listener.ConcurrentMessageListenerContainer;
import org.springframework.kafka.listener.ContainerProperties;
import org.springframework.kafka.listener.MessageListener;
import org.springframework.kafka.support.serializer.ErrorHandlingDeserializer;
import org.springframework.kafka.support.serializer.StringDeserializer;

import java.util.Collections;
import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;

import static org.junit.jupiter.api.Assertions.assertTrue;

public class KafkaContainerTest {

    // 封装测试用的监听器,用CountDownLatch验证消息接收
    static class TestableMessageListener implements MessageListener<String, String> {
        private final CountDownLatch latch;
        private final StringBuilder receivedMessages = new StringBuilder();

        public TestableMessageListener(int expectedMsgCount) {
            this.latch = new CountDownLatch(expectedMsgCount);
        }

        @Override
        public void onMessage(ConsumerRecord<String, String> record) {
            receivedMessages.append(record.value()).append(",");
            latch.countDown();
        }

        public boolean awaitCompletion(long timeout, TimeUnit unit) throws InterruptedException {
            return latch.await(timeout, unit);
        }

        public String getReceivedMessages() {
            return receivedMessages.toString();
        }
    }

    @Test
    void testConsumeMultipleMessages() throws InterruptedException {
        // 1. 配置消费者必要属性
        Map<String, Object> consumerConfigs = new HashMap<>();
        consumerConfigs.put("group.id", "test-group");
        consumerConfigs.put("key.deserializer", ErrorHandlingDeserializer.class);
        consumerConfigs.put("value.deserializer", ErrorHandlingDeserializer.class);
        consumerConfigs.put(ErrorHandlingDeserializer.KEY_DESERIALIZER_CLASS, StringDeserializer.class.getName());
        consumerConfigs.put(ErrorHandlingDeserializer.VALUE_DESERIALIZER_CLASS, StringDeserializer.class.getName());

        // 2. 创建并配置MockConsumer
        MockConsumer<String, String> mockConsumer = new MockConsumer<>(OffsetResetStrategy.EARLIEST);
        TopicPartition topicPartition = new TopicPartition("my-topic", 0);
        mockConsumer.assign(Collections.singleton(topicPartition));
        mockConsumer.seek(topicPartition, 0);

        // 3. 自定义ConsumerFactory返回MockConsumer
        org.springframework.kafka.core.ConsumerFactory<String, String> mockConsumerFactory = new org.springframework.kafka.core.ConsumerFactory<>() {
            @Override
            public org.apache.kafka.clients.consumer.Consumer<String, String> createConsumer() {
                return mockConsumer;
            }

            @Override
            public org.apache.kafka.clients.consumer.Consumer<String, String> createConsumer(String groupId, String clientIdPrefix) {
                return createConsumer();
            }

            @Override
            public Map<String, Object> getConfigurationProperties() {
                return consumerConfigs;
            }
        };

        // 4. 初始化测试监听器,预期接收3条消息
        TestableMessageListener testListener = new TestableMessageListener(3);

        // 5. 构建容器
        ContainerProperties containerProperties = new ContainerProperties("my-topic");
        containerProperties.setMessageListener(testListener);
        ConcurrentMessageListenerContainer<String, String> container = new ConcurrentMessageListenerContainer<>(
                mockConsumerFactory,
                containerProperties
        );

        // 6. 添加模拟消息
        mockConsumer.addRecord(new ConsumerRecord<>("my-topic", 0, 0L, "key1", "message1"));
        mockConsumer.addRecord(new ConsumerRecord<>("my-topic", 0, 1L, "key2", "message2"));
        mockConsumer.addRecord(new ConsumerRecord<>("my-topic", 0, 2L, "key3", "message3"));

        try {
            // 7. 启动容器并等待处理完成
            container.start();
            boolean allProcessed = testListener.awaitCompletion(5, TimeUnit.SECONDS);
            assertTrue(allProcessed, "Timeout waiting for messages to be consumed");
            // 可额外验证消息内容
            assertTrue(testListener.getReceivedMessages().contains("message1,message2,message3"));
        } finally {
            // 8. 停止容器避免线程泄漏
            container.stop();
        }
    }
}

Kotlin/Kotest 示例

采用Kotest行为驱动风格实现:

import org.apache.kafka.clients.consumer.ConsumerRecord
import org.apache.kafka.clients.consumer.MockConsumer
import org.apache.kafka.clients.consumer.OffsetResetStrategy
import org.apache.kafka.common.TopicPartition
import org.springframework.kafka.listener.ConcurrentMessageListenerContainer
import org.springframework.kafka.listener.ContainerProperties
import org.springframework.kafka.listener.MessageListener
import org.springframework.kafka.support.serializer.ErrorHandlingDeserializer
import org.springframework.kafka.support.serializer.StringDeserializer
import io.kotest.core.spec.style.BehaviorSpec
import io.kotest.matchers.shouldBe
import io.kotest.matchers.string.shouldContain
import java.util.Collections
import java.util.concurrent.CountDownLatch
import java.util.concurrent.TimeUnit

class KafkaContainerTest : BehaviorSpec({

    given("a configured Kafka message listener container") {
        val targetTopic = "my-topic"
        val expectedMessageCount = 3
        val completionLatch = CountDownLatch(expectedMessageCount)
        val receivedMessages = mutableListOf<String>()

        // 模拟业务监听器逻辑
        val testListener = MessageListener<String, String> { record ->
            receivedMessages.add(record.value())
            completionLatch.countDown()
        }

        // 创建并配置MockConsumer
        val mockConsumer = MockConsumer<String, String>(OffsetResetStrategy.EARLIEST).apply {
            val topicPartition = TopicPartition(targetTopic, 0)
            assign(Collections.singleton(topicPartition))
            seek(topicPartition, 0)
            // 添加模拟消息
            addRecord(ConsumerRecord(targetTopic, 0, 0L, "key1", "msg1"))
            addRecord(ConsumerRecord(targetTopic, 0, 1L, "key2", "msg2"))
            addRecord(ConsumerRecord(targetTopic, 0, 2L, "key3", "msg3"))
        }

        // 自定义ConsumerFactory
        val mockConsumerFactory = object : org.springframework.kafka.core.ConsumerFactory<String, String> {
            override fun createConsumer(): org.apache.kafka.clients.consumer.Consumer<String, String> = mockConsumer
            override fun createConsumer(groupId: String?, clientIdPrefix: String?): org.apache.kafka.clients.consumer.Consumer<String, String> = createConsumer()
            override fun getConfigurationProperties() = mapOf(
                "group.id" to "kotlin-test-group",
                "key.deserializer" to ErrorHandlingDeserializer::class.java,
                "value.deserializer" to ErrorHandlingDeserializer::class.java,
                ErrorHandlingDeserializer.KEY_DESERIALIZER_CLASS to StringDeserializer::class.java.name,
                ErrorHandlingDeserializer.VALUE_DESERIALIZER_CLASS to StringDeserializer::class.java.name
            )
        }

        // 构建容器
        val container = ConcurrentMessageListenerContainer(
            mockConsumerFactory,
            ContainerProperties(targetTopic).apply {
                messageListener = testListener
            }
        )

        `when`("the container is started") {
            container.start()
            val completed = completionLatch.await(5, TimeUnit.SECONDS)

            then("all messages should be consumed and processed") {
                completed shouldBe true
                receivedMessages.size shouldBe expectedMessageCount
                receivedMessages.joinToString(",") shouldContain "msg1,msg2,msg3"
            }
        }

        afterTest {
            container.stop()
        }
    }
})

关键注意事项

  • 分区绑定:必须给MockConsumer调用assign()指定分区或subscribe()订阅主题,否则容器无法发现可消费的消息
  • 序列化配置:使用ErrorHandlingDeserializer包裹真实反序列化器,避免模拟过程中出现序列化异常
  • 容器生命周期:测试完成后务必停止容器,防止线程泄漏
  • 异步验证:用CountDownLatch或线程安全集合跟踪消息处理状态,确保异步消费完成后再执行断言

内容的提问来源于stack exchange,提问作者k314159

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 17:23:15