如何为生成多matplotlib图表的函数编写适配pytest-mpl的测试用例
我有一个基于matplotlib实现的my_plots()函数,可创建并展示多个图表。我希望使用pytest-mpl插件对这些图表进行测试,但该插件要求每个测试用例仅对应单个图表。
我当前的实现方案是:通过monkeypatch编写get_figures()方法生成图表序列,再通过pytest.mark.parametrize将图表传入测试函数。该方案可正常完成图像比对逻辑,但由于原始绘图函数是在测试用例外部执行的,无法被pytest.mark.filterwarnings这类pytest内置特性,或是自定义fixture所作用。
当前实现代码:
import warnings import matplotlib.pyplot as plt import pytest def my_plots(): plt.figure() plt.plot([0, 1]) plt.show() plt.figure() warnings.warn("bad things may happen") plt.plot([0, 2]) plt.show() def get_figures(): def no_show(): pass with pytest.MonkeyPatch.context() as mp: mp.setattr(plt, "show", no_show) my_plots() for fig_num in plt.get_fignums(): yield plt.figure(fig_num) @pytest.mark.filterwarnings("error") @pytest.mark.mpl_image_compare @pytest.mark.parametrize("fig", get_figures()) def test_my_plots(fig): return fig
原写法失效的核心是pytest.mark.parametrize的参数解析发生在测试收集阶段,此时测试运行的上下文还未初始化:测试函数上挂载的警告过滤规则、自定义fixture、monkeypatch作用域都是测试执行阶段才会生效的,收集阶段运行get_figures()时这些规则完全没有加载,自然无法作用到绘图逻辑上。
另外收集阶段生成并持有matplotlib的figure实例,很容易污染全局状态,导致不同测试用例的图表互相干扰。
核心思路是把实际执行绘图的逻辑从测试收集阶段移到测试执行阶段,通过indirect参数化配合fixture承接绘图逻辑:
基础实现
import warnings import matplotlib.pyplot as plt import pytest def my_plots(): plt.figure() plt.plot([0, 1]) plt.show() plt.figure() warnings.warn("bad things may happen") plt.plot([0, 2]) plt.show() @pytest.fixture def fig(request): fig_idx = request.param created_figs = [] def capture_fig(): # 拦截plt.show,捕获当前生成的图表 created_figs.append(plt.gcf()) # 同组测试仅执行一次绘图逻辑,避免重复运行拖慢速度 if fig_idx == 0: mp = pytest.MonkeyPatch() mp.setattr(plt, "show", capture_fig) request.addfinalizer(mp.undo) my_plots() # 测试结束后自动清理所有图表,避免全局状态污染 request.addfinalizer(lambda: [plt.close(f) for f in created_figs]) return created_figs[fig_idx] @pytest.mark.filterwarnings("error") @pytest.mark.mpl_image_compare @pytest.mark.parametrize("fig", [0, 1], indirect=True) def test_my_plots(fig): return fig
这个实现的特点:
- 绘图逻辑运行在测试执行阶段,
filterwarnings、自定义fixture、monkeypatch都能正常生效 - 同一组测试仅运行一次
my_plots(),没有重复执行的性能损耗 - 自动清理matplotlib全局缓存,不会出现图表串扰问题
自动识别图表数量
如果不想手动指定图表索引,可以加pytest收集钩子,自动统计my_plots()生成的图表总数,不需要硬编码参数列表:
def pytest_generate_tests(metafunc): if "fig" not in metafunc.fixturenames: return # 收集阶段仅统计图表数量,不保留实例 def dummy_show(): pass with pytest.MonkeyPatch.context() as mp: mp.setattr(plt, "show", dummy_show) my_plots() total_figs = len(plt.get_fignums()) plt.close("all") # 动态生成参数,实际绘图逻辑仍在fixture中执行 metafunc.parametrize("fig", range(total_figs), indirect=True)
加上这段代码后,原测试函数上的@pytest.mark.parametrize装饰器就可以删掉,测试会自动适配my_plots()生成的任意数量图表。
注意:收集阶段只做数量统计即可,千万不要在这个阶段持有figure实例,否则很容易出现内存泄漏、图表状态异常的问题。
内容的提问来源于stack exchange,提问作者RuthC

