如何为Python中计算嵌套深度的递归函数编写类型提示以通过Pyright严格模式检查?
如何为Python中计算嵌套深度的递归函数编写类型提示以通过Pyright严格模式检查?
我完全理解你的困扰——你想要一个能处理任意嵌套结构的通用递归函数,但Pyright严格模式下的类型检查一直报错,那些换成具体类型(比如把Any改成int|str、Mapping换成dict)的临时方案又破坏了函数的通用性,确实让人头疼。
报错原因分析
Pyright之所以报“类型部分未知”的错误,核心问题在于你原来的Nested类型定义用了Any,再加上递归的联合类型(Any | Mapping[str, Nested] | Iterable[Nested]),让Pyright无法明确推断遍历Mapping.values()或Iterable元素时的具体类型——Any会模糊类型边界,导致递归调用时的参数类型无法被精准确认。
解决方案:泛型递归类型定义
我们可以通过泛型递归类型替代原来的Any,让类型边界更清晰,同时保留函数的通用性,完美通过Pyright严格模式检查:
# pyright: strict from collections.abc import Iterable, Mapping from typing import TypeVar, Union # 定义一个TypeVar,代表所有非嵌套的基础类型(比如int、str、bool等) T = TypeVar('T') # 递归定义嵌套类型:要么是基础类型T,要么是字符串键的嵌套Mapping,要么是非字符串的嵌套Iterable Nested = Union[T, Mapping[str, 'Nested[T]'], Iterable['Nested[T]']] def _get_max_depth(obj: Nested[T]) -> int: if isinstance(obj, Mapping): return max([0] + [_get_max_depth(val) for val in obj.values()]) + 1 elif isinstance(obj, Iterable) and not isinstance(obj, str): return max([0] + [_get_max_depth(elt) for elt in obj]) + 1 else: return 0
这个写法的关键是用TypeVar明确了“基础类型”的边界,让Pyright能精准区分:哪些是不需要递归的基础值,哪些是需要继续遍历的嵌套结构,彻底消除了类型推断的模糊性。
备选方案:用Protocol定义嵌套结构接口
如果你不想用泛型,也可以通过Protocol来定义嵌套结构的行为,明确哪些类型可以被递归处理,同样能通过严格模式检查:
# pyright: strict from collections.abc import Iterable, Mapping, Protocol from typing import Union, Any # 定义一个Protocol,描述可递归遍历的嵌套结构 class NestedStructure(Protocol): def __getitem__(self, key: str) -> 'NestedStructure': ... def __iter__(self) -> Iterable['NestedStructure']: ... # 嵌套类型要么是任意基础类型,要么是符合NestedStructure的结构 Nested = Union[Any, NestedStructure] def _get_max_depth(obj: Nested) -> int: if isinstance(obj, Mapping): return max([0] + [_get_max_depth(val) for val in obj.values()]) + 1 elif isinstance(obj, Iterable) and not isinstance(obj, str): return max([0] + [_get_max_depth(elt) for elt in obj]) + 1 else: return 0
这种写法通过Protocol约定了嵌套结构的接口,既保持了函数的通用性,又让Pyright能清晰识别可递归处理的类型。
备注:内容来源于stack exchange,提问作者Marcin Barczyński
相关产品推荐
相关产品推荐

