如何Mock Python函数以避免其在模块导入时被调用?
问题:模块导入时执行的函数抛出异常,无法测试目标函数(不可修改原代码)
原代码(无法修改)
app/annoying_file.py:
def annoying_function(): '''Does something that generates exception due to some hardcoded cloud stuff''' raise ValueError() # Simulate the original function raising error due to no cloud connection annoying_variable = annoying_function() def normal_function(): '''Works fine by itself''' return True
初始测试代码
def test_normal_function(): from app.annoying_file import normal_function assert normal_function() == True
测试失败的错误堆栈
failed: def test_normal_function(): > from app.annoying_file import normal_function test\test_annoying_file.py:6: _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ app\annoying_file.py:6: in <module> annoying_variable = annoying_function() _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ def annoying_function(): '''Does something that generates exception due to some hardcoded cloud stuff''' > raise ValueError() E ValueError app\annoying_file.py:3: ValueError
尝试的Mock方案及失败
尝试通过pytest-mock提前patch函数,但依然失败:
def test_normal_function(mocker): mocker.patch("app.annoying_file.annoying_function", return_value="foo") from app.annoying_file import normal_function assert normal_function() == True
错误堆栈:
failed: thing = <module 'app' (<_frozen_importlib_external._NamespaceLoader object at 0x00000244A7C72FE0>)> comp = 'annoying_file', import_path = 'app.annoying_file' def _dot_lookup(thing, comp, import_path): try: > return getattr(thing, comp) E AttributeError: module 'app' has no attribute 'annoying_file' ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1238: AttributeError During handling of the above exception, another exception occurred: mocker = <pytest_mock.plugin.MockerFixture object at 0x00000244A7C72380> def test_normal_function(mocker): > mocker.patch("app.annoying_file.annoying_function", return_value="foo") test\test_annoying_file.py:5: _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ .venv\lib\site-packages\pytest_mock\plugin.py:440: in __call__ return self._start_patch( .venv\lib\site-packages\pytest_mock\plugin.py:258: in _start_patch mocked: MockType = p.start() ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1585: in start result = self.__enter__() ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1421: in __enter__ self.target = self.getter() ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1608: in <lambda> getter = lambda: _importer(target) ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1251: in _importer thing = _dot_lookup(thing, comp, import_path) ....\.pyenv\pyenv-win\versions\3.10.5\lib\unittest\mock.py:1240: in _dot_lookup __import__(import_path) app\annoying_file.py:6: in <module> annoying_variable = annoying_function() _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ _ def annoying_function(): '''Does something that generates exception due to some hardcoded cloud stuff''' > raise ValueError() E ValueError app\annoying_file.py:3: ValueError
失败原因:patch操作需要先导入目标模块,而模块导入时会立即执行顶层的annoying_variable = annoying_function(),导致异常提前触发,patch无法生效。
解决方案
方案一:手动加载模块并提前注入Mock
利用importlib.util手动加载模块,在执行模块代码前替换掉annoying_function,避免顶层代码抛出异常:
import importlib.util import sys from unittest.mock import Mock def test_normal_function(): # 模块路径和名称 module_path = "app/annoying_file.py" module_name = "app.annoying_file" # 1. 获取模块spec spec = importlib.util.spec_from_file_location(module_name, module_path) # 2. 创建空模块对象 module = importlib.util.module_from_spec(spec) # 3. 提前注入mock的annoying_function到模块中 module.annoying_function = Mock(return_value="mocked_value") # 4. 注册模块到sys.modules,避免重复导入 sys.modules[module_name] = module try: # 5. 执行模块代码(此时调用的是mock函数,不会抛出异常) spec.loader.exec_module(module) # 6. 测试目标函数 assert module.normal_function() == True finally: # 清理sys.modules,不影响其他测试 del sys.modules[module_name]
方案二:使用mock.patch强制创建属性拦截
通过patch的create=True参数,在模块导入前动态创建annoying_function属性:
from unittest.mock import patch def test_normal_function(): # 在导入模块前patch,create=True允许创建不存在的模块属性 with patch("app.annoying_file.annoying_function", return_value="foo", create=True): from app.annoying_file import normal_function assert normal_function() == True
原理:patch会在模块加载时,自动给app.annoying_file模块添加mock的annoying_function,顶层代码执行时调用的是mock函数,不会触发异常。
内容的提问来源于stack exchange,提问作者TheArturro
相关产品推荐
相关产品推荐

