Python输入类型检查装饰器:如何提取list注解的子类型?
提取函数注解中list子类型的解决方法
当你使用list[str]这类泛型注解时,它在Python中属于types.GenericAlias类型,要提取其中的子类型,有两种可靠的方法:
方法一:直接访问泛型的内部属性
GenericAlias类型实例有__origin__和__args__两个关键属性:
__origin__:返回泛型的原始类型(比如list[str].__origin__返回list)__args__:返回泛型的参数元组(比如list[str].__args__返回(str,))
以下是实现类型检查装饰器的示例代码:
import types from functools import wraps def type_check(func): @wraps(func) def wrapper(*args, **kwargs): annotations = func.__annotations__ param_names = list(annotations.keys()) for idx, (name, anno_type) in enumerate(annotations.items()): if name == 'return': continue # 获取对应参数的值 if name in kwargs: value = kwargs[name] else: value = args[idx] # 处理list泛型的情况 if isinstance(anno_type, types.GenericAlias) and anno_type.__origin__ == list: elem_type = anno_type.__args__[0] if not all(isinstance(item, elem_type) for item in value): raise TypeError(f"参数 {name} 的所有元素必须是 {elem_type.__name__} 类型") elif not isinstance(value, anno_type): raise TypeError(f"参数 {name} 必须是 {anno_type.__name__} 类型") return func(*args, **kwargs) return wrapper # 测试 @type_check def process_str_list(items: list[str]): print("处理完成:", items) process_str_list(["apple", "banana"]) # 正常执行 process_str_list([1, 2, "3"]) # 抛出TypeError
方法二:使用标准库的typing.get_args()和typing.get_origin()
官方推荐用这两个函数处理泛型类型,它们对Python不同版本的泛型写法(比如3.9+的list[str]和旧版本的typing.List[str])都能兼容,且不依赖内部属性:
from functools import wraps from typing import get_args, get_origin def type_check(func): @wraps(func) def wrapper(*args, **kwargs): annotations = func.__annotations__ param_names = list(annotations.keys()) for idx, (name, anno_type) in enumerate(annotations.items()): if name == 'return': continue # 获取参数值 value = kwargs.get(name, args[idx]) # 判断是否为list泛型 origin = get_origin(anno_type) if origin == list: elem_type = get_args(anno_type)[0] if not all(isinstance(item, elem_type) for item in value): raise TypeError(f"参数 {name} 的元素必须是 {elem_type.__name__} 类型") elif not isinstance(value, anno_type): raise TypeError(f"参数 {name} 必须是 {anno_type.__name__} 类型") return func(*args, **kwargs) return wrapper # 测试 @type_check def process_int_list(items: list[int]): print("处理完成:", items) process_int_list([1, 2, 3]) # 正常执行 process_int_list([1, "2", 3]) # 抛出TypeError
额外说明
- 对于嵌套泛型(比如
list[list[int]]),可以通过递归调用的方式逐层提取子类型。 - 如果你需要支持更多泛型类型(如
dict、tuple),只需扩展origin的判断条件即可,逻辑和处理list一致。
内容的提问来源于stack exchange,提问作者luka hasch
相关产品推荐
相关产品推荐

