如何在pytest中配置跨测试文件共享的公共Spark Session
问题根因
你遇到的会话冲突、列访问KeyError问题核心来自两个配置错误:
- Spark 2.0及以上版本中,
SQLContext/HiveContext的能力已经完全集成到SparkSession中,手动基于SparkContext重复创建SQLContext实例时,不同实例的parquet元数据缓存、临时视图注册表是独立的,跨测试调用时会出现元数据不同步,直接导致列名找不到的KeyError。 - 你在测试文件中写的
get_data是普通函数,没有被pytest识别为fixture,依赖注入逻辑不生效,部分场景下会拿到未正确初始化的上下文实例,加剧会话冲突。
可行修复方案(pytest单Spark Session跨文件共享最佳实践)
1. 修正conftest.py的session级Spark fixture
把项目根目录的conftest.py配置调整为如下内容,保证整个测试周期只有一个SparkSession实例,清理逻辑更严谨:
import os import shutil import pytest import pyspark from pyspark.sql import SparkSession @pytest.fixture(scope="session") def spark_session(): # 测试场景固定用本地模式,根据CPU核数设置并行度 conf = ( pyspark.SparkConf() .setMaster("local[*]") .setAppName("pytest-spark-test") .set("spark.sql.shuffle.partitions", "1") # 测试场景降低并行度提速 .set("spark.driver.memory", "2g") .set("spark.ui.enabled", "false") # 测试时关掉UI减少资源占用 ) spark = ( SparkSession.builder .config(conf=conf) .getOrCreate() ) # 启动前先清一遍残留缓存,避免历史运行的元数据干扰 spark.catalog.clearCache() yield spark # teardown逻辑:先停会话再删临时目录,避免文件占用报错 spark.stop() if os.path.exists("./spark-warehouse"): shutil.rmtree("./spark-warehouse", ignore_errors=True) if os.path.exists("./metastore_db"): shutil.rmtree("./metastore_db", ignore_errors=True)
2. 调整测试文件的读数据逻辑
不要手动创建SQLContext,直接使用注入的spark_session的内置read接口,建议把通用读数据逻辑也封装成fixture,按需设置scope:
import pytest import pandas as pd # 如果所有测试用例都用同一份数据,可以把scope设为session,只读一次提速 @pytest.fixture(scope="session") def test_data(spark_session): # 直接用SparkSession自带的read API,走统一元数据上下文 df = spark_session.read.parquet("data/expected.parquet") # 如果你需要转pandas DataFrame,直接在这里做转换即可,不用每个测试重复写 pandas_df = df.toPandas() return pandas_df # 测试用例直接注入fixture即可 def test_data_column_exist(test_data): # 这里访问列就不会再抛KeyError assert "target_column" in test_data.columns
3. 额外注意事项
- 所有Spark相关操作统一使用注入的
spark_session实例提供的API,不要手动新建SQLContext、SparkContext实例,避免上下文不一致。 - 如果单个测试需要创建临时视图、临时表,可以新增一个function级的fixture做自动清理,避免污染全局会话,示例代码:
@pytest.fixture(scope="function") def clean_temp_view(spark_session): yield # 每个测试跑完删掉所有自定义临时视图 for view in spark_session.catalog.listTables("default"): if view.isTemporary: spark_session.catalog.dropTempView(view.name) - 如果还是出现列相关报错,可以在读取数据后加一行
df.printSchema()打印schema,确认parquet文件本身的列名没有拼写错误、特殊字符问题。
内容的提问来源于stack exchange,提问作者Mistapopo
相关产品推荐
相关产品推荐

