如何为接收任意大小构造器元组的Python函数添加类型提示?
问题
我可以编写如下函数:
async def par_all(fns: tuple[Awaitable[T], Awaitable[U]]) -> tuple[T,U]: ...
但如何将其扩展为接收任意大小的元组,且不丢失类型信息?(理想情况不用*fns,但如果需要也可以接受)
我尝试了如下写法,但在pyright中无法正常工作:
from typing import Awaitable, TypeVarTuple, Unpack Ts = TypeVarTuple('Ts') async def par_all(fns: tuple[Awaitable[Unpack[Ts]], ...]) -> tuple[Unpack[Ts]]: return await asyncio.gather(*fns) import asyncio async def foo() -> int: await asyncio.sleep(1) return 1 async def bar() -> str: await asyncio.sleep(1) return "hello" async def baz() -> float: await asyncio.sleep(1) return 3.14 async def main(): result = await par_all((foo(), bar(), baz())) print(result) # (1, "hello", 3.14) # result的类型应为tuple[int, str, float],但未能实现! asyncio.run(main())
另外,若要接收Iterable[Awaitable[T]]并返回tuple[T],最佳方式是不是使用重载?
解决方案
一、处理任意大小元组并保留类型信息
你之前的写法问题在于元组类型标注的语法逻辑错误,正确的做法是利用TypeVarTuple和Unpack的组合,直接标注元组每个元素为对应类型的Awaitable。以下是Pyright可正确识别的写法:
from typing import Awaitable, TypeVarTuple, Unpack import asyncio Ts = TypeVarTuple('Ts') async def par_all(fns: tuple[Awaitable[Ts], ...]) -> tuple[Unpack[Ts]]: return await asyncio.gather(*fns) async def foo() -> int: await asyncio.sleep(1) return 1 async def bar() -> str: await asyncio.sleep(1) return "hello" async def baz() -> float: await asyncio.sleep(1) return 3.14 async def main(): result = await par_all((foo(), bar(), baz())) reveal_type(result) # Pyright会正确输出: tuple[int, str, float] print(result) asyncio.run(main())
这里Ts作为TypeVarTuple会被展开为元组每个位置的具体类型,Awaitable[Ts]会对应每个元素是包裹对应类型的Awaitable,Pyright能准确推断返回值的元组类型。
二、处理Iterable[Awaitable[T]]的情况
如果要接收Iterable[Awaitable[T]]并返回tuple[T],最简洁实用的方式是直接用泛型标注:
from typing import Awaitable, Iterable, TypeVar, Tuple import asyncio T = TypeVar('T') async def par_all_iterable(fns: Iterable[Awaitable[T]]) -> Tuple[T, ...]: return await asyncio.gather(*fns)
这种写法能让Pyright推断返回值为tuple[T, ...],其中T是所有Awaitable返回类型的公共超类型。如果需要针对固定长度的可迭代对象(比如确定元素数量的列表)做精确类型推断,再考虑使用重载:
from typing import Awaitable, Iterable, TypeVar, Tuple, overload, List T = TypeVar('T') U = TypeVar('U') V = TypeVar('V') @overload async def par_all_iterable(fns: List[Awaitable[T]]) -> Tuple[T, ...]: ... @overload async def par_all_iterable(fns: List[Awaitable[T], Awaitable[U]]) -> Tuple[T, U]: ... @overload async def par_all_iterable(fns: List[Awaitable[T], Awaitable[U], Awaitable[V]]) -> Tuple[T, U, V]: ... async def par_all_iterable(fns: Iterable[Awaitable[T]]) -> Tuple[T, ...]: return await asyncio.gather(*fns)
不过这种重载需要提前定义固定长度的场景,对于任意长度的可迭代对象,泛型标注的方式适用性更强。
内容的提问来源于stack exchange,提问作者xyzzyrz
相关产品推荐
相关产品推荐

