如何编写PySpark Kafka写入流测试用例?遇AttributeError问题求助
错误原因分析
你遇到的AttributeError是因为patch的目标对象完全错误:
- 业务代码中,
start()方法返回的是pyspark.sql.streaming.StreamingQuery实例,awaitTermination()是这个实例的方法 - 但你的测试代码错误地patch了
pyspark.streaming.StreamingContext.awaitTermination,而StreamingContext类本身并没有这个属性,因此触发报错。
修复步骤与完善后的测试用例
1. 修正Patch目标
将针对StreamingContext的patch替换为针对StreamingQuery的patch。
2. 调整Mock参数顺序
Python unittest.mock.patch的参数传递是逆序的:装饰器从上到下,测试函数参数从后到前对应。
3. 设置正确的Mock返回值
- 链式调用的
DataStreamWriter每个方法需返回自身,保证调用链不中断 start()方法需返回带有awaitTermination方法的StreamingQuerymock对象
4. 补充缺失的Patch
测试中用到的pyspark_option_mock等对应DataStreamWriter.option方法,需要补充对应的patch。
以下是修改后的完整测试用例:
from unittest import mock import pyspark.sql.streaming import kafka_utils # 替换为你的业务模块名 class TestKafkaStreamWriter: @mock.patch("kafka_utils.get_secrets") # 补充get_secrets的patch @mock.patch("pyspark.sql.streaming.StreamingQuery.awaitTermination") @mock.patch("pyspark.sql.streaming.DataStreamWriter.start") @mock.patch("pyspark.sql.streaming.DataStreamWriter.trigger") @mock.patch("pyspark.sql.streaming.DataStreamWriter.option") @mock.patch("pyspark.sql.streaming.DataStreamWriter.format") @mock.patch("pyspark.sql.streaming.DataStreamWriter.outputMode") @mock.patch("pyspark.sql.DataFrame.writeStream") def test_kafka_writestream( self, write_stream_mock, output_mode_mock, format_mock, option_mock, trigger_mock, start_mock, await_termination_mock, get_secrets_mock, spark_session ): # 模拟get_secrets返回值 get_secrets_mock.return_value = {"username": "test_api_key", "password": "test_secret"} # 创建mock的DataStreamWriter,链式调用返回自身 mock_data_stream_writer = mock.Mock(spec=pyspark.sql.streaming.DataStreamWriter) output_mode_mock.return_value = mock_data_stream_writer format_mock.return_value = mock_data_stream_writer option_mock.return_value = mock_data_stream_writer trigger_mock.return_value = mock_data_stream_writer # 创建mock的StreamingQuery,包含awaitTermination方法 mock_streaming_query = mock.Mock(spec=pyspark.sql.streaming.StreamingQuery) start_mock.return_value = mock_streaming_query # 调用业务函数 df_stream = spark_session.createDataFrame([(1, "test")], ["id", "value"]) SECRETS_ARN = "test_arn" SECRETS_REGION = "test_region" KAFKA_BOOTSTRAP_SERVER = "test_bootstrap" checkpoint_location = "test_file_path" target_file_path = "target_file_path" KAFKA_TOPIC = "test_topic" kafka_utils.streamwriter( df_stream, SECRETS_ARN, SECRETS_REGION, KAFKA_BOOTSTRAP_SERVER, checkpoint_location, target_file_path, KAFKA_TOPIC, ) # 验证调用 write_stream_mock.assert_called_once() output_mode_mock.assert_called_once_with("append") format_mock.assert_called_once_with("kafka") option_mock.assert_has_calls( [ mock.call("header", True), mock.call( "kafka.sasl.jaas.config", "org.apache.kafka.common.security.plain.PlainLoginModule required username=test_api_key password=test_secret", ), mock.call("kafka.bootstrap.servers", KAFKA_BOOTSTRAP_SERVER), mock.call("kafka.sasl.mechanism", "PLAIN"), mock.call("kafka.security.protocol", "SASL_SSL"), mock.call("checkpointLocation", f"{checkpoint_location}/{target_file_path}"), mock.call("topic", KAFKA_TOPIC), ], any_order=False # 确保调用顺序正确 ) trigger_mock.assert_called_once_with(processingTime="10 seconds") start_mock.assert_called_once() await_termination_mock.assert_called_once() get_secrets_mock.assert_called_once_with(SECRETS_ARN, SECRETS_REGION)
额外说明
- 用
spec参数创建mock对象,能保证mock的行为更接近真实对象,避免误测试 - 链式调用的mock必须返回自身,否则会因为调用链中断导致属性不存在错误
- 验证调用时添加
any_order=False确保参数顺序和业务代码一致
内容的提问来源于stack exchange,提问作者Perin
相关产品推荐
相关产品推荐

