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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 04:35:01