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)
解决方案
核心修改思路
- 移除自定义
AwaitableMock,直接用AsyncMock的原生特性处理可等待对象; - 让
asyncio.as_completed的Mock直接返回传入的任务列表,保证await task能获取到get_repository的返回值; - 通过
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
相关产品推荐
相关产品推荐

