Python3.11以下版本如何用unittest断言类型注解?
针对Python<3.11的类型断言方案(基于unittest)
核心思路
直接对比List[str]和type(dummy())会失败,因为type(dummy())返回的是运行时的list类型,而List[str]是静态泛型注解,并非运行时类型。要高效完成校验,需拆分容器类型和元素类型的检查逻辑,同时避免大数据集下的全量遍历。
具体实现
1. 提取函数返回类型注解
通过函数的__annotations__属性获取定义好的返回类型注解:
from typing import List import unittest def dummy() -> List[str]: return ["a", "b", "c"] class TestTypeAssertions(unittest.TestCase): def test_dummy_return_type(self): # 获取函数的返回类型注解 return_annotation = dummy.__annotations__["return"] result = dummy()
2. 容器与元素类型校验
利用Python3.8+官方提供的typing.get_origin和typing.get_args工具解析泛型类型,配合抽样检查元素类型:
from typing import get_origin, get_args def test_dummy_return_type(self): return_annotation = dummy.__annotations__["return"] result = dummy() # 校验容器的原始类型(比如List对应的list) container_type = get_origin(return_annotation) or return_annotation self.assertEqual(type(result), container_type) # 校验元素类型:抽样检查而非全量遍历 element_type = get_args(return_annotation)[0] if result: # 取前10个元素抽样,可根据需求调整样本量 sample_size = min(10, len(result)) for elem in result[:sample_size]: self.assertIsInstance(elem, element_type) else: # 空容器单独断言类型合规 self.assertIsInstance(result, container_type)
3. 封装通用断言方法
将逻辑封装成TestCase的扩展方法,方便复用:
from typing import get_origin, get_args class TypeAssertingTestCase(unittest.TestCase): def assert_return_type(self, func, expected_type, sample_size=10): result = func() # 处理非泛型类型(比如直接返回str的情况) origin_type = get_origin(expected_type) or expected_type self.assertIsInstance(result, origin_type) # 处理泛型容器类型 args = get_args(expected_type) if args: element_type = args[0] if result: sample = result[:min(sample_size, len(result))] for elem in sample: self.assertIsInstance(elem, element_type) # 使用示例 class TestDummy(TypeAssertingTestCase): def test_dummy(self): self.assert_return_type(dummy, List[str])
高效性说明
- 抽样检查:大数据集下仅校验部分元素,大幅减少校验耗时,同时常规场景下足以检测类型异常(若元素类型不一致,抽样大概率能覆盖)。
- 官方工具依赖:
get_origin和get_args是Python3.8+官方提供的泛型解析工具,比手动解析内部属性更可靠。
兼容低版本Python(3.7及以下)
如果使用Python3.7,可手动替代get_origin和get_args:
def get_origin_compat(t): return getattr(t, "__origin__", t) def get_args_compat(t): return getattr(t, "__args__", ())
只需将代码中的get_origin和get_args替换为上述兼容函数即可。
内容的提问来源于stack exchange,提问作者alphazwest
相关产品推荐
相关产品推荐

