如何根据所选辅助函数向主函数传递特定附加参数?
动态传递辅助函数参数并提示用户所需参数的解决方案
问题背景
我正在开发一款分析任务类库:
aux类包含32个辅助函数,所有函数都共享相同的基础输入参数,同时各自拥有专属的自定义参数Main类及同类类中有超100个主函数,用户通过calculation_tool参数指定计算所用的辅助函数(已通过字典映射实现关联)
由于辅助函数与主函数数量庞大,若在主函数中罗列所有辅助函数的参数会导致代码冗长混乱,需解决两个核心问题:
- 根据用户所选的辅助函数,自动传递对应所需的参数
- 明确告知用户当前选择下需要提供的特定参数
代码示例
辅助类与主类结构
class aux: def foo1(self, a, b, c): return 'apple' def foo2(self, a, b, c, var1, var2): return 'apple' def foo3(self, a, b, c, var1, var2, var3): return 'apple' def foo4(self, a, b, c, var1, var4, var5): return 'apple' def foo5(self, a, b, c, var1, var2, var4, var5): return 'apple' class Main: def func1(self, a, b, c, calculation_tool: int): return 'tomato' def func2(self, a, b, c, calculation_tool: int): return 'tomato'
主函数调用示例
def on_balance_volume( self, price_df: PandasDataFrame, n: int = 5, input_mode: int = 2, calculation_tool: int = 0, ) -> PandasDataFrame: """ :param price_df: Dataframe that contains price data from which Accumulation Distribution Indicator will be calculated. :param n: Lookback period of On Balance Volume indicator. :param input_mode: Defines from what kind of data MA is calculated. :param calculation_tool: Defines MA function from which OBV is calculated. :return: DataFrame with calculated OBV. """ _util = Utils() function, name = _util.choose_ma(calculation_tool) _change = _util.change(price_df=price_df, input_mode=input_mode) sign = _util.signum(price_df=_change, from_price=False, indicator_name="Change") product = pd.DataFrame() product["product"] = sign["Signum(Change)"] * price_df["Volume"] product_sum = pd.DataFrame() product_sum["On Balance Volume"] = product["product"].cumsum() obv_sml = function( price_df=product_sum, n=n, from_price=False, indicator_name="On Balance Volume" ) obv_sml.rename(columns={f"{name}{n}": f"On Balance Volume SmL {n}"}, inplace=True) obv = pd.concat( [product_sum.reset_index(drop=True), obv_sml.reset_index(drop=True)], axis=1 ) return obv
解决方案
1. 用可变关键字参数动态传递参数
主函数通过**kwargs接收用户传入的自定义参数,直接传递给选中的辅助函数,避免在主函数中硬编码所有辅助函数的参数。
修改后主类示例:
class Main: def func1(self, a, b, c, calculation_tool: int, **kwargs): aux_instance = aux() # 已有的工具映射字典 tool_map = { 1: aux_instance.foo1, 2: aux_instance.foo2, 3: aux_instance.foo3, 4: aux_instance.foo4, 5: aux_instance.foo5 } selected_func = tool_map[calculation_tool] # 传递基础参数+用户传入的专属参数 return selected_func(a, b, c, **kwargs)
2. 自动提取并校验专属参数
利用inspect模块解析辅助函数的签名,过滤出专属参数(排除基础参数),实现参数校验与提示。
参数提取工具函数
import inspect def get_custom_params(func, base_params=['a', 'b', 'c']): sig = inspect.signature(func) all_params = list(sig.parameters.keys()) # 过滤基础参数,得到当前辅助函数的专属参数 return [param for param in all_params if param not in base_params]
整合校验逻辑到主函数
class Main: def func1(self, a, b, c, calculation_tool: int, **kwargs): aux_instance = aux() tool_map = { 1: (aux_instance.foo1, "foo1"), 2: (aux_instance.foo2, "foo2"), 3: (aux_instance.foo3, "foo3"), 4: (aux_instance.foo4, "foo4"), 5: (aux_instance.foo5, "foo5") } selected_func, func_name = tool_map[calculation_tool] custom_params = get_custom_params(selected_func) # 校验用户是否传入了所有必需的专属参数 missing_params = [p for p in custom_params if p not in kwargs] if missing_params: raise ValueError(f"选择{func_name}需提供以下专属参数: {', '.join(missing_params)}") return selected_func(a, b, c, **kwargs)
3. 优化主函数调用示例(以OBV为例)
def on_balance_volume( self, price_df: PandasDataFrame, n: int = 5, input_mode: int = 2, calculation_tool: int = 0, **kwargs, ) -> PandasDataFrame: """ :param price_df: 用于计算累积派发指标的价格数据DataFrame。 :param n: 平衡成交量指标的回溯周期。 :param input_mode: 定义计算移动平均线的数据类型。 :param calculation_tool: 定义计算OBV所用的移动平均线函数。 :param kwargs: 所选移动平均线函数的专属参数,具体参数取决于calculation_tool的选择。 :return: 包含计算后OBV的DataFrame。 """ _util = Utils() function, name = _util.choose_ma(calculation_tool) # 提取当前MA函数的专属参数(排除固定传递的基础参数) base_ma_params = ['price_df', 'n', 'from_price', 'indicator_name'] custom_params = get_custom_params(function, base_ma_params) # 参数校验 missing_params = [p for p in custom_params if p not in kwargs] if missing_params: raise ValueError(f"选择{name}需提供以下专属参数: {', '.join(missing_params)}") _change = _util.change(price_df=price_df, input_mode=input_mode) sign = _util.signum(price_df=_change, from_price=False, indicator_name="Change") product = pd.DataFrame() product["product"] = sign["Signum(Change)"] * price_df["Volume"] product_sum = pd.DataFrame() product_sum["On Balance Volume"] = product["product"].cumsum() # 传递基础参数+专属参数 obv_sml = function( price_df=product_sum, n=n, from_price=False, indicator_name="On Balance Volume", **kwargs ) obv_sml.rename(columns={f"{name}{n}": f"On Balance Volume SmL {n}"}, inplace=True) obv = pd.concat( [product_sum.reset_index(drop=True), obv_sml.reset_index(drop=True)], axis=1 ) return obv
内容的提问来源于stack exchange,提问作者Jakub Szurlej
相关产品推荐
相关产品推荐

