堆叠装饰器时如何保留函数参数信息以支持反射检查?
堆叠装饰器时保留函数参数反射信息的解决方法
问题背景
需要实现一个装饰器,根据函数传入的service参数值自动注入对应的查询语句。当前实现的装饰器在单个使用时正常,但堆叠多个装饰器后,inspect.getfullargspec无法获取原函数的参数信息,导致逻辑失效。
现有装饰器代码:
import functools import inspect from typing import Optional def pass_query(service: str, query: str): def decorate(decoratee): @functools.wraps(decoratee) def decorator(*args, **kwargs): params = inspect.getfullargspec(decoratee) if "query" in params.args and service == params.get("service", None): kwargs["query"] = query return decoratee(*args, **kwargs) return decorator return decorate
堆叠装饰器的失效场景:
@pass_query("foo", "q1.sql") @pass_query("bar", "q2.sql") def track_report(service: str, query: Optional[str] = None) -> None: pass
问题原因
堆叠装饰器时,decoratee参数指向的是上层装饰器返回的包装函数,而非最底层的原函数。尽管functools.wraps会复制原函数的元数据,但inspect.getfullargspec默认不会递归解析__wrapped__属性(该属性由functools.wraps添加,指向被包装的函数),因此无法获取原函数的参数定义。
解决方案
1. 递归获取原函数
实现工具函数,通过__wrapped__属性递归追溯到最底层的原函数:
def get_original_func(func): while hasattr(func, "__wrapped__"): func = func.__wrapped__ return func
2. 修改装饰器逻辑
将原装饰器中获取参数信息的逻辑替换为获取原函数的参数,同时修正service参数的判断逻辑(原代码错误地从参数定义中取service值,实际应从调用时的args/kwargs中提取):
import functools import inspect from typing import Optional def get_original_func(func): while hasattr(func, "__wrapped__"): func = func.__wrapped__ return func def pass_query(service: str, query: str): def decorate(decoratee): @functools.wraps(decoratee) def decorator(*args, **kwargs): # 获取原函数的参数定义 original_func = get_original_func(decoratee) params = inspect.getfullargspec(original_func) # 提取当前调用的service参数值 service_arg_idx = params.args.index("service") current_service = kwargs.get("service") if current_service is None and len(args) > service_arg_idx: current_service = args[service_arg_idx] # 匹配service时注入query if current_service == service: kwargs["query"] = query return decoratee(*args, **kwargs) return decorator return decorate
3. 验证效果
使用堆叠装饰器测试:
@pass_query("foo", "q1.sql") @pass_query("bar", "q2.sql") def track_report(service: str, query: Optional[str] = None) -> None: print(f"Service: {service}, Query: {query}") track_report("foo") # 输出: Service: foo, Query: q1.sql track_report("bar") # 输出: Service: bar, Query: q2.sql track_report("test") # 输出: Service: test, Query: None
关键说明
functools.wraps会自动为包装函数添加__wrapped__属性,通过递归遍历该属性可以准确获取原函数。- 原代码中
params.get("service")是逻辑错误,params存储的是函数的参数名称列表,而非调用时传入的参数值,必须从args或kwargs中提取实际传入的service值。
内容的提问来源于stack exchange,提问作者t3chb0t
相关产品推荐
相关产品推荐

