如何避免Python装饰器包装函数重复定义参数?
避免装饰器wrapper重复定义参数的Pythonic方案
问题场景
你编写了Spark日期转换函数standard_date_formatter,并实现了参数校验装饰器validate_old_new,但当前装饰器的wrapper函数重复定义了与原函数完全一致的参数列表,希望找到更优雅的方式避免这种重复。
解决方案
1. 使用*args和**kwargs接收参数
通过通用参数接收所有输入,再手动解析出所需变量,彻底避免重复定义参数列表:
from functools import wraps import pyspark.sql.functions as f from pyspark.sql.types import DateType, StringType from typing import Optional from pyspark.sql import DataFrame def validate_old_new(fn): @wraps(fn) def wrapper(*args, **kwargs): # 解析参数:优先取关键字参数,无则按位置取 df = kwargs.get('df') or args[0] prev_name = kwargs.get('prev_name') or args[1] prev_format = kwargs.get('prev_fmt') or args[2] next_name = kwargs.get('next_name') or (args[3] if len(args) >=4 else None) next_format = kwargs.get('next_fmt') or (args[4] if len(args) >=5 else "yyyy-MM-dd") # 原校验逻辑 try: prev_type = df.schema[prev_name].dataType except KeyError: raise ValueError("Column to be converted does not exist.") if prev_type not in (StringType(), DateType()): raise AttributeError( "Column to be converted must be StringType or DateType" ) if not next_name: next_name = prev_name df_to_format = ( df if prev_type == StringType() else df.withColumn(prev_name, f.col(prev_name).cast(StringType())) ) return fn(df_to_format, prev_name, prev_format, next_name, next_format) return wrapper @validate_old_new def standard_date_formatter( df: DataFrame, prev_name: str, prev_fmt: str, next_name: Optional[str] = None, next_fmt: str = "yyyy-MM-dd" ): # 数据转换实现逻辑 pass
优缺点:实现简单,无需依赖额外模块;但参数解析依赖位置顺序,可读性稍差,原函数参数顺序变更时需同步调整解析逻辑。
2. 用inspect模块解析原函数签名
通过inspect模块自动获取原函数的参数签名,绑定并解析输入参数,无需手动处理位置顺序:
from functools import wraps import pyspark.sql.functions as f from pyspark.sql.types import DateType, StringType from typing import Optional from pyspark.sql import DataFrame import inspect def validate_old_new(fn): @wraps(fn) def wrapper(*args, **kwargs): # 获取原函数签名并绑定输入参数 sig = inspect.signature(fn) bound_args = sig.bind(*args, **kwargs) bound_args.apply_defaults() # 自动填充默认参数 # 提取所需变量 df = bound_args.arguments['df'] prev_name = bound_args.arguments['prev_name'] prev_format = bound_args.arguments['prev_fmt'] next_name = bound_args.arguments['next_name'] next_format = bound_args.arguments['next_fmt'] # 原校验逻辑 try: prev_type = df.schema[prev_name].dataType except KeyError: raise ValueError("Column to be converted does not exist.") if prev_type not in (StringType(), DateType()): raise AttributeError( "Column to be converted must be StringType or DateType" ) if not next_name: next_name = prev_name df_to_format = ( df if prev_type == StringType() else df.withColumn(prev_name, f.col(prev_name).cast(StringType())) ) # 更新参数后调用原函数 bound_args.arguments['df'] = df_to_format bound_args.arguments['next_name'] = next_name return fn(*bound_args.args, **bound_args.kwargs) return wrapper @validate_old_new def standard_date_formatter( df: DataFrame, prev_name: str, prev_fmt: str, next_name: Optional[str] = None, next_fmt: str = "yyyy-MM-dd" ): # 数据转换实现逻辑 pass
优缺点:自动适配原函数的参数结构,无需关心位置顺序,可读性和健壮性更强;需要依赖inspect模块,但属于Python标准库,无需额外安装。
3. 用数据类统一参数结构
通过dataclass封装所有参数,让原函数和装饰器都基于这个数据类操作,从根源上避免参数重复定义:
from dataclasses import dataclass from typing import Optional from pyspark.sql import DataFrame from functools import wraps import pyspark.sql.functions as f from pyspark.sql.types import DateType, StringType # 定义统一的参数数据类 @dataclass class DateFormatterParams: df: DataFrame prev_name: str prev_fmt: str next_name: Optional[str] = None next_fmt: str = "yyyy-MM-dd" def validate_old_new(fn): @wraps(fn) def wrapper(params: DateFormatterParams): # 原校验逻辑 try: prev_type = params.df.schema[params.prev_name].dataType except KeyError: raise ValueError("Column to be converted does not exist.") if prev_type not in (StringType(), DateType()): raise AttributeError( "Column to be converted must be StringType or DateType" ) processed_next_name = params.next_name if params.next_name else params.prev_name processed_df = ( params.df if prev_type == StringType() else params.df.withColumn(params.prev_name, f.col(params.prev_name).cast(StringType())) ) # 创建处理后的参数实例 processed_params = DateFormatterParams( df=processed_df, prev_name=params.prev_name, prev_fmt=params.prev_fmt, next_name=processed_next_name, next_fmt=params.next_fmt ) return fn(processed_params) return wrapper @validate_old_new def standard_date_formatter(params: DateFormatterParams): # 数据转换实现逻辑,直接使用params的属性 pass # 调用示例 # params = DateFormatterParams(df=my_spark_df, prev_name="raw_date", prev_fmt="dd/MM/yyyy") # standard_date_formatter(params)
优缺点:参数结构清晰,新增/修改参数只需修改数据类,所有使用处自动同步;适合参数较多、后续可能变更的场景,是最符合Pythonic风格的方案。
内容的提问来源于stack exchange,提问作者oogway74
相关产品推荐
相关产品推荐

