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

如何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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 18:57:26