如何将柯里化函数的最后位置参数传入剩余的末位参数?
问题
给定一个基于toolz的柯里化函数:
import toolz @toolz.curry def mklist(x, y, z): return [x, y, z]
目前mklist(x=1, y=2)(z=3)可以正常调用,但希望支持mklist(x=1, y=2)(3)的调用形式,让位置参数3自动绑定到最后一个参数z,而不是触发TypeError: mklist() got multiple values for argument 'x'错误。
解决方案
要实现这种调用逻辑,核心是让后续传入的位置参数自动填充到未被绑定的剩余参数的末尾,可以通过自定义柯里化装饰器或包装toolz的柯里化函数来实现。
方法1:自定义柯里化装饰器
这个装饰器会跟踪已绑定的参数,当后续传入位置参数时,自动将其分配给未指定的末位参数:
import inspect def curry_last(func): sig = inspect.signature(func) param_names = list(sig.parameters.keys()) def wrapper(*args, **kwargs): # 确定已绑定的参数 bound = sig.bind_partial(*args, **kwargs) bound.apply_defaults() used_params = set(bound.arguments.keys()) remaining_params = [p for p in param_names if p not in used_params] def inner(*inner_args, **inner_kwargs): new_kwargs = {**kwargs, **inner_kwargs} # 将位置参数从后往前匹配剩余参数 for param, arg in zip(reversed(remaining_params), reversed(inner_args)): if param not in new_kwargs: new_kwargs[param] = arg # 调用原函数 return func(*args, **new_kwargs) # 如果所有参数已绑定,直接返回结果 if not remaining_params: return func(*args, **kwargs) return inner return wrapper # 使用自定义装饰器 @curry_last def mklist(x, y, z): return [x, y, z] # 测试调用 print(mklist(x=1, y=2)(3)) # 输出: [1, 2, 3] print(mklist(1)(y=2)(3)) # 输出: [1, 2, 3] print(mklist(1, 2)(z=3)) # 输出: [1, 2, 3]
方法2:包装toolz的柯里化函数
如果想保留toolz.curry的原有特性,可以在柯里化后的函数外层加一层包装,处理位置参数的分配逻辑:
import toolz import inspect @toolz.curry def _mklist(x, y, z): return [x, y, z] def mklist(*args, **kwargs): sig = inspect.signature(_mklist) param_names = list(sig.parameters.keys()) def wrapped(*inner_args, **inner_kwargs): # 获取第一次调用后已绑定的参数 bound = sig.bind_partial(*args, **kwargs) used_params = set(bound.arguments.keys()) remaining_params = [p for p in param_names if p not in used_params] # 将位置参数分配给剩余的末位参数 new_inner_kwargs = {**inner_kwargs} for param, arg in zip(reversed(remaining_params), reversed(inner_args)): if param not in new_inner_kwargs: new_inner_kwargs[param] = arg return _mklist(*args, **kwargs)(**new_inner_kwargs) # 尝试直接调用,如果参数齐全则返回结果 try: return _mklist(*args, **kwargs) except TypeError: return wrapped # 测试 print(mklist(x=1, y=2)(3)) # 输出: [1, 2, 3] print(mklist(1, 2)(3)) # 输出: [1, 2, 3]
内容的提问来源于stack exchange,提问作者lupl
相关产品推荐
相关产品推荐

