如何在Python中Mock SparkSession.table?测试遇列不可迭代错误
问题分析与修复方案
1. 修复expected_df的语法错误
你的测试代码中创建expected_df时存在语法错误,导致Spark误将列名列表当作数据行处理,触发「Column not iterable」错误:
# 错误写法:列名被混入数据列表,缺少闭合括号 expected_df = spark.createDataFrame([(1, 100, "ketchup"),["id", "qty", "condiment"])
修正为正确的参数格式:
# 数据列表与列名列表是两个独立参数 expected_df = spark.createDataFrame([(1, 100, "ketchup")], ["id", "qty", "condiment"])
2. 修正Mock路径
你当前Mock的是SparkSession类,但实际需要Mock的是main_file中实际使用的sparkSession实例的table方法。假设main_file中使用全局sparkSession对象,调整Mock装饰器:
# 替换原Mock路径,指向main_file中的sparkSession实例的table方法 @mock.patch("main_file.sparkSession.table")
如果sparkSession是Class类的实例属性(如self.sparkSession),则需调整为对应实例属性的Mock路径。
3. 修正断言逻辑
你直接断言类实例obj等于expected_df是错误的——obj是类实例,expected_df是DataFrame。应捕获transform_table的返回值,再比较DataFrame内容:
# 获取函数返回结果 result_df = obj.transform_table("useless table name", 100, 150) # 转换为列表后比较(DataFrame不能直接用==判断内容) assert result_df.collect() == expected_df.collect()
4. 修复原函数的column变量问题
原函数中column.isBetween未明确指定列,需补充导入并指定过滤列(比如qty列):
# main_file.py中补充导入 from pyspark.sql.functions import col def transform_table(table_name, start, end): # 替换为具体列的过滤逻辑 return sparkSession.table(table_name).filter(col("qty").isBetween(start, end))
完整修复后的测试代码
@mock.patch("main_file.sparkSession.table") @pytest.mark.usefixtures("spark") def test_transform_table(self, mocked_table, spark): injected_df = spark.createDataFrame( [(1, 100, "ketchup"), (2, 200, "mayo")], ["id", "qty", "condiment"] ) expected_df = spark.createDataFrame( [(1, 100, "ketchup")], ["id", "qty", "condiment"] ) mocked_table.return_value = injected_df obj = Class(spark) result_df = obj.transform_table("useless table name", 100, 150) assert result_df.collect() == expected_df.collect()
内容的提问来源于stack exchange,提问作者user3746406
相关产品推荐
相关产品推荐

