高阶函数类型注解方案:如何让mypy根据原函数推断新函数的重复参数签名
实现符合预期的类型注解方案
要让mypy能自动推断出duplicate返回函数的签名,你需要用到Python 3.10+ typing模块里的ParamSpec和Concatenate特性,它们能帮我们精准捕获原函数的参数信息,再组合成新函数的参数规格。
完整实现代码
from typing import Callable, TypeVar, ParamSpec, Concatenate import inspect # 定义泛型:捕获原函数的返回类型 R = TypeVar('R') # 定义参数规格泛型:捕获原函数的参数类型信息 P = ParamSpec('P') def duplicate(f: Callable[P, R]) -> Callable[Concatenate[P, P], tuple[R, R]]: def g(*args: P.args, **kwargs: P.kwargs) -> tuple[R, R]: # 获取原函数的签名信息,用来拆分参数 sig = inspect.signature(f) param_list = list(sig.parameters.values()) num_pos_params = 0 for param in param_list: # 统计常规位置/关键字参数的数量 if param.kind in (param.POSITIONAL_ONLY, param.POSITIONAL_OR_KEYWORD): num_pos_params += 1 # 暂不支持带*args或关键字-only参数的函数 elif param.kind == param.VAR_POSITIONAL: raise ValueError("duplicate doesn't support functions with *args") elif param.kind == param.KEYWORD_ONLY: raise ValueError("duplicate doesn't support functions with keyword-only parameters") # 检查传入的参数数量是否符合预期(原参数数量的两倍) if len(args) != 2 * num_pos_params: raise TypeError(f"Expected {2 * num_pos_params} positional arguments, got {len(args)}") # 拆分参数为两组,分别传给原函数两次调用 args1 = args[:num_pos_params] args2 = args[num_pos_params:] return (f(*args1), f(*args2)) return g
代码说明
泛型与参数规格:
TypeVar('R'):用来捕获原函数的返回类型,确保新函数返回的元组元素类型和原函数一致。ParamSpec('P'):捕获原函数的完整参数类型信息(比如f(a:int, b:str)的P就是(int, str))。Concatenate[P, P]:把原函数的参数规格重复两次,直接定义了新函数g的参数签名——也就是你想要的“原参数两次重复”。
运行时参数处理:
我们用inspect模块获取原函数的签名,统计常规位置参数的数量,这样就能把g收到的参数拆分成两组,分别传给原函数调用。目前这个实现暂不支持带*args或关键字-only参数的函数,这类场景会导致参数拆分逻辑失效。
测试示例
def f(a: int, b: str) -> float: return float(a) + len(b) # mypy会自动推断g的类型为:(int, str, int, str) -> tuple[float, float] g = duplicate(f) result = g(1, "hello", 2, "world") print(result) # 输出:(6.0, 7.0)
这个方案完全符合你的需求,mypy能准确识别g的参数和返回值类型,运行时也能正确处理参数拆分。
内容的提问来源于stack exchange,提问作者Chris Grimm
相关产品推荐
相关产品推荐

