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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 19:43:09