如何Mock导入的scan_all函数?单元测试断言失败求助
无法Mock
scan_all 函数的问题 我尝试了多种方法都没得到预期结果,很多相关问题用@patch装饰器做Mock,但我想用更简化的方式。现在无法Mock特定的scan_all函数,测试时断言始终失败(左侧为空列表,右侧是预期的Job对象),且scan_all从未被调用。
主类代码
from databaseClass import scan_all class JobsDB: def __init__(self, jobs_table): self.jobs_table = jobs_table async def get_all_jobs(self, include_executions=False): response = await scan_all(self.jobs_table) if include_executions: return [JobWithExecutions(**item) for item in response] return [Job(**item) for item in response]
数据库类代码
async def scan_all(table, **kwargs): response = await table.scan(**kwargs) items = response['Items'] while 'LastEvaluatedKey' in response: response = await table.scan(**kwargs, ExclusiveStartKey=response['LastEvaluatedKey']) items = items + response['Items'] return items
测试类代码
import pytest from unittest.mock import AsyncMock, MagicMock from db import JobsDB, Job @pytest.fixture def mock_scan_all(): return AsyncMock() @pytest.fixture def jobs_table(): return AsyncMock() @pytest.fixture def jobs_db(): return JobsDB(AsyncMock()) @pytest.fixture def job(): return Job( id='1839e18a-898b-42d1-8747-9b5495dbb0a6', status='PENDING', description='Test', start_date='2024-01-22T13:00:00.000Z', end_date='2024-01-22T14:00:00.000Z', ) @pytest.mark.database @pytest.mark.asyncio async def test_get_all_jobs(jobs_db, jobs_table, mock_scan_all, job): # This used a mock_scan_all fixture in the parameters mock_scan_all.return_value = [vars(job)] jobs_table.scan_all = mock_scan_all # Call the get_all_jobs method result = await jobs_db.get_all_jobs() # jobs_db.scan_all.assert_called_once_with(jobs_db.jobs_table) # Assert that the result is a list containing the expected Job object assert result == [job]
我尝试过的方法
# Didn't work # jobs_db.jobs_table.scan_all = AsyncMock(return_value=[vars(job)]) # Didn't work # jobs_db.scan_all.return_value = [vars(job)] # Didn't work # scan_all = AsyncMock() # scan_all.return_value = [vars(job)] # Didn't work # mock = AsyncMock() # mock.scan_all.return_value = [vars(job)] # Didn't work # mock_scan_all = AsyncMock() # mock_scan_all.return_value = [vars(job)] # jobs_table = MagicMock() # jobs_table.scan_all = mock_scan_all # Didn't work # jobs_table.scan.return_value = {'Items': [vars(job) for job in jobs]}
问题原因及解决方案
问题出在你Mock的对象不对:JobsDB类里是直接从databaseClass导入并调用scan_all函数,而不是调用jobs_table.scan_all方法。所以你之前给jobs_table添加scan_all属性的操作完全没用,代码里根本没这么调用。
要解决这个问题,你需要Mock的是JobsDB所在模块(也就是db模块)中导入的scan_all对象,具体修改测试代码如下:
@pytest.mark.database @pytest.mark.asyncio async def test_get_all_jobs(jobs_db, mock_scan_all, job): # 关键:替换db模块里的scan_all为mock对象 import db db.scan_all = mock_scan_all mock_scan_all.return_value = [vars(job)] # 调用方法 result = await jobs_db.get_all_jobs() # 验证scan_all被正确调用 mock_scan_all.assert_called_once_with(jobs_db.jobs_table) # 断言结果 assert result == [job]
或者用monkeypatch fixture来替换(更简洁的方式):
@pytest.mark.database @pytest.mark.asyncio async def test_get_all_jobs(jobs_db, mock_scan_all, job, monkeypatch): monkeypatch.setattr("db.scan_all", mock_scan_all) mock_scan_all.return_value = [vars(job)] result = await jobs_db.get_all_jobs() mock_scan_all.assert_called_once_with(jobs_db.jobs_table) assert result == [job]
这样就能正确Mockscan_all函数,让它返回你预设的值,测试就能通过了。
内容的提问来源于stack exchange,提问作者Darius Fiallo
相关产品推荐
相关产品推荐

