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

高阶函数类型注解方案:如何让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

代码说明

  1. 泛型与参数规格:

    • TypeVar('R'):用来捕获原函数的返回类型,确保新函数返回的元组元素类型和原函数一致。
    • ParamSpec('P'):捕获原函数的完整参数类型信息(比如f(a:int, b:str)的P就是(int, str))。
    • Concatenate[P, P]:把原函数的参数规格重复两次,直接定义了新函数g的参数签名——也就是你想要的“原参数两次重复”。
  2. 运行时参数处理:
    我们用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.27 15:44:05