Python中为保留结构的嵌套映射函数添加类型注解
嘿,我之前刚好为了实现一个能严格保留原嵌套结构的映射函数,折腾了好一会儿类型注解的事儿——既要让函数逻辑能正确处理dict、NamedTuple和各种可迭代对象,又要让mypy这类类型检查工具能准确校验输入输出的类型,确实踩了几个小坑。下面给你一步步拆解:
首先,先把核心的函数逻辑写出来,这个函数要能识别并保留原结构:
from collections.abc import Iterable, Mapping from typing import Callable, TypeVar, Any A = TypeVar('A') B = TypeVar('B') def nested_map(f: Callable[[A], B], data: Any) -> Any: # 处理字典:保留原字典类型,递归映射每个值 if isinstance(data, Mapping): return type(data)({k: nested_map(f, v) for k, v in data.items()}) # 处理NamedTuple:通过_asdict属性判断是否是NamedTuple,递归处理每个字段后生成原类型的实例 elif isinstance(data, tuple) and hasattr(data, '_asdict'): field_values = (nested_map(f, getattr(data, field)) for field in data._fields) return type(data)(*field_values) # 处理其他可迭代对象,但排除字符串和字节串(它们是可迭代但我们要当叶子节点) elif isinstance(data, Iterable) and not isinstance(data, (str, bytes)): return type(data)(nested_map(f, item) for item in data) # 叶子节点:直接应用传入的函数f else: return f(data)
这段逻辑跑起来没问题,但用Any当输入输出类型的话,类型检查工具完全帮不上忙——比如你不小心把处理字符串的函数传给了全是整数的嵌套结构,mypy根本不会提醒你,这就失去了类型注解的意义。
接下来就是关键的类型注解优化。我们需要定义一个递归的泛型类型别名,用来描述“任意嵌套的可迭代/映射结构,叶子节点是指定类型”。Python 3.10+支持用字符串引用的递归类型,写法如下:
from typing import TypeVar, Callable, Iterable, Mapping, Tuple, NamedTuple, Union, TypeAlias A = TypeVar('A') B = TypeVar('B') # 定义递归类型:Nested[A] 可以是A本身,或者嵌套的Iterable/ Mapping/ 元组,每个元素/值都是Nested[A] Nested: TypeAlias = Union[ A, Iterable['Nested[A]'], Mapping[Any, 'Nested[A]'], Tuple['Nested[A]', ...] ]
然后把函数的类型注解替换成这个递归泛型:
def nested_map(f: Callable[[A], B], data: Nested[A]) -> Nested[B]: if isinstance(data, Mapping): return type(data)({k: nested_map(f, v) for k, v in data.items()}) # type: ignore elif isinstance(data, tuple) and hasattr(data, '_asdict'): field_values = (nested_map(f, getattr(data, field)) for field in data._fields) return type(data)(*field_values) # type: ignore elif isinstance(data, Iterable) and not isinstance(data, (str, bytes)): return type(data)(nested_map(f, item) for item in data) # type: ignore else: return f(data)
这里的# type: ignore是因为类型检查工具没法自动推断出type(data)(比如list、dict或者自定义NamedTuple)的泛型参数——比如当data是list[Nested[A]]时,type(data)是list,我们返回的是list[Nested[B]],逻辑上完全正确,但mypy没法识别type(data)的泛型参数变化,所以需要暂时忽略这个小错误。
现在,类型检查工具就能帮我们做很多事了:比如如果你的输入结构的叶子节点类型和f的参数类型不匹配,mypy会直接报错;如果你的输出结构和预期的嵌套类型不符,也会给出明确提示。
举个实际的测试例子:
# 定义一个同构的NamedTuple(所有字段都是int类型,或者嵌套结构) class Product(NamedTuple): id: int price: int # 构造测试输入 test_input = { "electronics": [Product(1, 1999), Product(2, 2999)], "clothing": (Product(3, 99), {"discount": Product(4, 49)}) } # 定义一个把int转成带货币符号的字符串的函数 def format_price(num: int) -> str: return f"¥{num}" # 调用nested_map test_output = nested_map(format_price, test_input)
这时候,test_output的结构和test_input完全一致,所有原来的int类型叶子(Product的id、price字段,还有嵌套结构里的int)都被转换成了带¥符号的字符串。mypy会正确推断出test_output的类型,并且如果你的函数f的参数类型和输入叶子类型不匹配,它会立刻提醒你。
最后再提个重要的细节:为什么要排除字符串和字节串?因为如果不排除的话,函数会把字符串拆成单个字符,然后对每个字符应用f,这显然不是我们想要的——我们希望字符串本身作为一个完整的叶子节点处理。
内容来源于stack exchange

