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

Python 3.12异步函数get_repositories的单元测试Mock问题求助

异步函数get_repositories单元测试Mock问题解决方案

问题描述

  • 无法正确Mock get_repository返回的协程对象,验证其是否传入asyncio.as_completed;
  • 无法让await task正确返回指定的Repository对象或None值。

待测试函数代码

async def get_repositories(
    gl_client: gitlab.client.Gitlab,
    marks: vng_release_notes_obj.Marks,
    tag_pattern: re.Pattern[str],
    repositories_info: Iterable[RepositoryBasicInfo],
) -> Optional[dict[pathlib.PurePosixPath, Repository]]:

    repositories: dict[pathlib.PurePosixPath, Repository] = {}
    semaphore = asyncio.Semaphore(4)
    tasks = [
        get_repository(semaphore, gl_client, marks, tag_pattern, repository_info)
        for repository_info in repositories_info
    ]
    for task in asyncio.as_completed(tasks):
        repository = await task

        if repository is None:
            return None

        repositories[repository.path] = repository
    return repositories

现有测试代码及问题

class TestGetRepositories(unittest.IsolatedAsyncioTestCase):

    def setUp(self) -> None:
        """Call before each test in this class."""
        self.last_modification = LastModification(
            "", datetime.datetime.now(), ""
        )
        self.gl_client = unittest.mock.AsyncMock()

        self.marks = unittest.mock.MagicMock(spec_set=Marks)

        self.tag_pattern = re.compile("dd")

        self.repository_info_a = RepositoryBasicInfo(
            "a", pathlib.PurePosixPath("a")
        )
        self.repository_info_b = RepositoryBasicInfo(
            "b", pathlib.PurePosixPath("b")
        )
        self.repository_info_c = RepositoryBasicInfo(
            "c", pathlib.PurePosixPath("c")
        )

        self.repository_info = [self.repository_info_a, self.repository_info_b, self.repository_info_c]

    class AwaitableMock(unittest.mock.AsyncMock):
        def __await__(self) -> typing.Iterator[typing.Any]:
            self.await_count += 1
            print(self.return_value)
            return iter([self.return_value.return_value])

    @unittest.mock.patch("asyncio.as_completed")
    @unittest.mock.patch("asyncio.Semaphore")
    @unittest.mock.patch("get_repository")
    async def test_invalid_case(
        self,
        mock_get_repository: unittest.mock.AsyncMock,
        mock_semaphore: unittest.mock.MagicMock,
        mock_as_completed: unittest.mock.MagicMock,
    ) -> None:
        """Test invalid case."""

        mock_semaphore_obj = unittest.mock.MagicMock()
        mock_semaphore.return_value = mock_semaphore_obj

        repository = unittest.mock.AsyncMock(spec=vng_release_notes.repository.Repository)
        repository.path = pathlib.PurePosixPath("a")

        #Here should return None in second variable
        as_compmleted_ret_a = self.AwaitableMock(return_value=repository)
        as_compmleted_ret_b = self.AwaitableMock(return_value=None)
        as_compmleted_ret_c = self.AwaitableMock()

        mock_as_completed.return_value = [as_compmleted_ret_a, as_compmleted_ret_b, as_compmleted_ret_c]
        self.assertIsNone(
            await vng_release_notes.gl.repository.get_repositories(
                self.gl_client, self.marks, self.tag_pattern, self.repository_info
            )
        )

        self.assertEqual(mock_get_repository.call_count, 3)

        self.assertEqual(len(mock_get_repository.call_args_list[0].args), 5)
        self.assertIs(mock_get_repository.call_args_list[0].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[0].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[0].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[0].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[0].args[4], self.repository_info_a)

        self.assertEqual(len(mock_get_repository.call_args_list[1].args), 5)
        self.assertIs(mock_get_repository.call_args_list[1].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[1].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[1].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[1].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[1].args[4], self.repository_info_b)

        self.assertEqual(len(mock_get_repository.call_args_list[1].args), 5)
        self.assertIs(mock_get_repository.call_args_list[2].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[2].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[2].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[2].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[2].args[4], self.repository_info_c)

        self.assertEqual(mock_semaphore.call_count, 1)
        self.assertIs(mock_semaphore.call_args[0][0], 4)

        self.assertEqual(mock_as_completed.call_count, 1)

    @unittest.mock.patch("asyncio.as_completed")
    @unittest.mock.patch("asyncio.Semaphore")
    @unittest.mock.patch(
        "get_repository", new_callable=unittest.mock.AsyncMock
    )
    async def test_valid_case(
        self,
        mock_get_repository: unittest.mock.AsyncMock,
        mock_semaphore: unittest.mock.MagicMock,
        mock_as_completed: unittest.mock.MagicMock,
    ) -> None:
        """Test valid case.

        Args:
            mock_get_repository: Mock for get_repository function.
            mock_semaphore: Mock for asyncio.Semaphore object.
            mock_as_completed: Mock for asyncio.as_completed function.
        """

        mock_semaphore_obj = unittest.mock.MagicMock()
        mock_semaphore.return_value = mock_semaphore_obj

        repository_a_path = pathlib.PurePosixPath("a")
        repository_a = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repository_a.path = repository_a_path

        repository_b_path = pathlib.PurePosixPath("b")
        repository_b = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repository_b.path = repository_b_path

        repository_c_path = pathlib.PurePosixPath("c")
        repository_c = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repository_c.path = repository_c_path

        as_compmleted_ret_a = self.AwaitableMock(side_effect=repository_a)
        as_compmleted_ret_b = self.AwaitableMock(side_effect=repository_b)
        as_compmleted_ret_c = self.AwaitableMock(side_effect=repository_c)

        mock_as_completed.return_value = [as_compmleted_ret_a, as_compmleted_ret_b, as_compmleted_ret_c]

        output = await vng_release_notes.gl.repository.get_repositories(
            self.gl_client, self.marks, self.tag_pattern, self.repository_info
        )

        if output is None:
            self.fail("Should not be None")

        self.assertIsInstance(output, dict)
        self.assertEqual(len(output), 3)

        self.assertIs(output[repository_a_path], repository_a)
        self.assertIs(output[repository_b_path], repository_b)
        self.assertIs(output[repository_c_path], repository_c)

        self.assertEqual(mock_get_repository.call_count, 3)

        self.assertEqual(len(mock_get_repository.call_args_list[0].args), 5)
        self.assertIs(mock_get_repository.call_args_list[0].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[0].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[0].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[0].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[0].args[4], self.repository_info_a)

        self.assertEqual(len(mock_get_repository.call_args_list[1].args), 5)
        self.assertIs(mock_get_repository.call_args_list[1].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[1].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[1].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[1].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[1].args[4], self.repository_info_b)

        self.assertEqual(len(mock_get_repository.call_args_list[1].args), 5)
        self.assertIs(mock_get_repository.call_args_list[2].args[0], mock_semaphore_obj)
        self.assertIs(mock_get_repository.call_args_list[2].args[1], self.gl_client)
        self.assertIs(mock_get_repository.call_args_list[2].args[2], self.marks)
        self.assertIs(mock_get_repository.call_args_list[2].args[3], self.tag_pattern)
        self.assertIs(mock_get_repository.call_args_list[2].args[4], self.repository_info_c)

        self.assertEqual(mock_semaphore.call_count, 1)
        self.assertIs(mock_semaphore.call_args[0][0], 4)

        self.assertEqual(mock_as_completed.call_count, 1)

解决方案

核心修改思路

  1. 移除自定义AwaitableMock,直接用AsyncMock的原生特性处理可等待对象;
  2. 让asyncio.as_completed的Mock直接返回传入的任务列表,保证await task能获取到get_repository的返回值;
  3. 通过mock_as_completed的调用参数验证任务是否正确传入。

修改后的测试代码示例

class TestGetRepositories(unittest.IsolatedAsyncioTestCase):

    def setUp(self) -> None:
        self.last_modification = LastModification("", datetime.datetime.now(), "")
        self.gl_client = unittest.mock.AsyncMock()
        self.marks = unittest.mock.MagicMock(spec_set=Marks)
        self.tag_pattern = re.compile("dd")
        
        self.repository_info_a = RepositoryBasicInfo("a", pathlib.PurePosixPath("a"))
        self.repository_info_b = RepositoryBasicInfo("b", pathlib.PurePosixPath("b"))
        self.repository_info_c = RepositoryBasicInfo("c", pathlib.PurePosixPath("c"))
        self.repository_info = [self.repository_info_a, self.repository_info_b, self.repository_info_c]

    @unittest.mock.patch("asyncio.as_completed")
    @unittest.mock.patch("asyncio.Semaphore")
    @unittest.mock.patch("get_repository")
    async def test_invalid_case(
        self,
        mock_get_repository: unittest.mock.AsyncMock,
        mock_semaphore: unittest.mock.MagicMock,
        mock_as_completed: unittest.mock.MagicMock,
    ) -> None:
        mock_semaphore_obj = unittest.mock.MagicMock()
        mock_semaphore.return_value = mock_semaphore_obj

        # 为每个get_repository调用设置不同的返回值
        repo_a = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repo_a.path = pathlib.PurePosixPath("a")
        mock_get_repository.side_effect = [
            unittest.mock.AsyncMock(return_value=repo_a),
            unittest.mock.AsyncMock(return_value=None),
            unittest.mock.AsyncMock()
        ]

        # 让as_completed返回传入的任务列表迭代器
        mock_as_completed.side_effect = lambda tasks, *args, **kwargs: iter(tasks)

        # 执行测试
        result = await vng_release_notes.gl.repository.get_repositories(
            self.gl_client, self.marks, self.tag_pattern, self.repository_info
        )
        self.assertIsNone(result)

        # 验证get_repository调用次数和参数
        self.assertEqual(mock_get_repository.call_count, 3)
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_a
        )
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_b
        )
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_c
        )

        # 验证任务是否正确传入as_completed
        self.assertEqual(mock_as_completed.call_count, 1)
        passed_tasks = mock_as_completed.call_args[0][0]
        self.assertEqual(len(passed_tasks), 3)
        for task, call in zip(passed_tasks, mock_get_repository.call_args_list):
            self.assertIs(task, call.return_value)

        # 验证Semaphore初始化参数
        mock_semaphore.assert_called_once_with(4)

    @unittest.mock.patch("asyncio.as_completed")
    @unittest.mock.patch("asyncio.Semaphore")
    @unittest.mock.patch("get_repository", new_callable=unittest.mock.AsyncMock)
    async def test_valid_case(
        self,
        mock_get_repository: unittest.mock.AsyncMock,
        mock_semaphore: unittest.mock.MagicMock,
        mock_as_completed: unittest.mock.MagicMock,
    ) -> None:
        mock_semaphore_obj = unittest.mock.MagicMock()
        mock_semaphore.return_value = mock_semaphore_obj

        # 创建三个Repository实例
        repo_a = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repo_a.path = pathlib.PurePosixPath("a")
        repo_b = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repo_b.path = pathlib.PurePosixPath("b")
        repo_c = unittest.mock.MagicMock(spec=vng_release_notes.repository.Repository)
        repo_c.path = pathlib.PurePosixPath("c")

        # 为每个get_repository调用设置返回值
        mock_get_repository.side_effect = [
            unittest.mock.AsyncMock(return_value=repo_a),
            unittest.mock.AsyncMock(return_value=repo_b),
            unittest.mock.AsyncMock(return_value=repo_c)
        ]

        # mock as_completed返回原任务列表
        mock_as_completed.side_effect = lambda tasks, *args, **kwargs: iter(tasks)

        # 执行测试
        result = await vng_release_notes.gl.repository.get_repositories(
            self.gl_client, self.marks, self.tag_pattern, self.repository_info
        )

        self.assertIsNotNone(result)
        self.assertIsInstance(result, dict)
        self.assertEqual(len(result), 3)
        self.assertIs(result[repo_a.path], repo_a)
        self.assertIs(result[repo_b.path], repo_b)
        self.assertIs(result[repo_c.path], repo_c)

        # 验证get_repository调用次数和参数
        self.assertEqual(mock_get_repository.call_count, 3)
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_a
        )
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_b
        )
        mock_get_repository.assert_any_call(
            mock_semaphore_obj, self.gl_client, self.marks, self.tag_pattern, self.repository_info_c
        )

        # 验证任务是否正确传入as_completed
        self.assertEqual(mock_as_completed.call_count, 1)
        passed_tasks = mock_as_completed.call_args[0][0]
        self.assertEqual(len(passed_tasks), 3)
        for task, call in zip(passed_tasks, mock_get_repository.call_args_list):
            self.assertIs(task, call.return_value)

        # 验证Semaphore
        mock_semaphore.assert_called_once_with(4)

关键修改点说明

  • 简化Mock对象:AsyncMock本身就是可等待对象,直接设置return_value就能控制await的结果,无需自定义__await__方法;
  • 模拟as_completed行为:让Mock返回传入的任务列表,保证测试逻辑和生产逻辑一致;
  • 验证任务传入:通过mock_as_completed.call_args获取传入的任务列表,对比每个任务是否是mock_get_repository返回的实例,确保参数传递正确;
  • 批量设置返回值:用side_effect为mock_get_repository的每次调用分配不同的返回值,模拟多任务的不同结果。

内容的提问来源于stack exchange,提问作者user25184856

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 11:49:57