如何为Python中参数数量未知的Callable编写泛型以匹配元组参数?
如何为自定义starmap生成器添加准确的类型提示?
我实现了一个等价于itertools.starmap的生成器,但当前的类型提示无法让func的参数类型与iterable中的元组类型对齐。原实现代码如下:
from typing import Callable, Generator, Iterable, Unpack, TypedDict def starmap[T](func: Callable[..., T], iterable: Iterable[tuple]) -> Generator[T, None, None]: yield from map(func, *zip(*iterable)) # 测试代码 import itertools args = lambda x, y: x + y, [(1, 2), (3, 4)] assert list(itertools.starmap(*args)) == list(starmap(*args))
我尝试引入第二个泛型Args并使用Unpack,但不确定是否是正确的做法,该如何编写泛型才能让参数类型对齐?
正确的泛型写法
你需要用一个泛型参数表示iterable中每个元组的类型,再通过Unpack将这个元组类型展开为func的参数列表,让类型检查器确保两者的类型匹配。修改后的代码如下:
from typing import Callable, Generator, Iterable, Unpack, TypeVar # 定义两个泛型参数:T代表返回值类型,Args代表参数元组类型 T = TypeVar('T') Args = TypeVar('Args', bound=tuple) def starmap[T, Args](func: Callable[[Unpack[Args]], T], iterable: Iterable[Args]) -> Generator[T, None, None]: yield from map(func, *zip(*iterable)) # 测试代码(类型检查器会正确识别参数类型) import itertools args = lambda x, y: x + y, [(1, 2), (3, 4)] assert list(itertools.starmap(*args)) == list(starmap(*args)) # 错误示例会被类型检查器捕获:元组长度不匹配func参数数量 # bad_args = lambda x, y: x + y, [(1, 2, 3), (4, 5, 6)] # list(starmap(*bad_args)) # 类型检查器会提示参数不匹配
代码说明
Args作为泛型参数,被限制为tuple的子类,用来表示iterable中每个元素的元组类型。Callable[[Unpack[Args]], T]表示func接受Args元组展开后的参数,并返回T类型的值。Iterable[Args]确保iterable中的每个元素都是符合func参数要求的元组,实现了类型对齐。
这样修改后,类型检查器就能准确校验func和iterable的类型兼容性,比如当iterable中的元组长度与func的参数数量不匹配时,会直接抛出类型错误。
内容的提问来源于stack exchange,提问作者Airat Valiullin
相关产品推荐
相关产品推荐

