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

pytest中如何让被测模块共享conftest.py的Spark Session并替换模块内的Spark方法?

pytest中如何让被测模块共享conftest.py的Spark Session并替换模块内的Spark方法?

嗨,我完全懂你的困扰!你在conftest.py里已经写了Spark Session的初始化函数,但被测模块里的函数还在调用自己那套没初始化好的Spark实例,导致测试卡壳对吧?别慌,咱们用pytest自带的工具就能搞定,下面给你一步步说清楚:

  • 第一步:先把conftest.py里的Session改成标准的pytest fixture
    你现在写的session()只是普通函数,pytest没法自动注入到测试里,得加上@pytest.fixture装饰器,还可以设置作用域提升测试效率,比如:

    import pytest
    from pyspark.sql import SparkSession
    
    @pytest.fixture(scope="session")  # 整个测试会话只用一个Spark实例,节省资源
    def spark_session():
        # 初始化Spark Session,按需配置参数
        spark = SparkSession.builder \
            .master("local[1]") \
            .appName("pytest-spark-test") \
            .getOrCreate()
        yield spark  # 把实例传给测试用例
        spark.stop()  # 测试结束后关闭会话
    
  • 第二步:用monkeypatch替换被测模块里的Spark对象
    假设你的file_to_be_tested.py里是直接用模块级的Spark实例,比如:

    # file_to_be_tested.py
    from pyspark.sql import SparkSession
    
    # 模块内的Spark实例,未正确初始化
    spark = SparkSession.builder.getOrCreate()
    
    def func1(df):
        # 内部调用了依赖这个spark的func2
        return func2(df)
    
    def func2(df):
        return spark.sql("SELECT * FROM df").count()
    

    那测试文件里就可以通过monkeypatch把模块里的spark替换成我们fixture里的实例:

    import file_to_be_tested as t
    
    def test_func(spark_session, monkeypatch):
        # 关键一步:把被测模块的spark替换成我们初始化好的spark_session
        monkeypatch.setattr(t, "spark", spark_session)
        
        # 准备测试用的DataFrame
        test_df = spark_session.createDataFrame([(1, "foo"), (2, "bar")], ["id", "name"])
        
        # 现在调用func1,内部的func2会自动用我们的Spark Session了
        result = t.func1(test_df)
        
        # 做你需要的断言
        assert result == 2
    
  • 如果被测模块是通过函数获取Spark的,比如:

    # file_to_be_tested.py
    def get_spark():
        return SparkSession.builder.getOrCreate()
    
    def func1(df):
        spark = get_spark()
        return spark.sql("...").count()
    

    那我们就替换这个get_spark函数,让它返回我们的fixture实例:

    def test_func(spark_session, monkeypatch):
        # 替换get_spark函数,直接返回我们的spark_session
        monkeypatch.setattr(t, "get_spark", lambda: spark_session)
        
        # 后续测试步骤和上面一样
        test_df = spark_session.createDataFrame([...])
        result = t.func1(test_df)
        assert result == xxx
    

简单来说,monkeypatch就是帮你在测试运行时“偷换”模块里的对象/函数,让被测代码完全使用我们在conftest里准备好的Spark Session,不用再担心模块自己初始化的无效实例了。你之前尝试的mock思路其实和这个是一致的,只是用pytest自带的monkeypatch会更贴合pytest的测试体系,用起来更顺手。

备注:内容来源于stack exchange,提问作者u09j

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 13:14:29