使用attrs库创建带互斥参数的Python类
解决互斥参数类的实例化问题
针对你遇到的这个问题,我有两种可行的解决方案,先从最彻底的方案说起,再讲保留灵活性的替代方式:
1. 完全禁用常规构造函数(推荐)
这是从根源上解决问题的办法,确保用户只能通过你定义的from_prices和from_returns方法来创建实例,彻底杜绝直接传入不兼容参数的可能。
具体实现步骤:
- 用
@attr.s(init=False)告诉attrs不要自动生成构造函数; - 自定义一个
__init__方法并抛出错误,阻止外部直接调用; - 在备选构造函数里用
cls.__new__(cls)创建实例(绕过__init__),然后手动给所有属性赋值。
修改后的代码如下:
import pandas as pd import attr @attr.s(init=False) # 关闭attrs自动生成构造函数 class MutuallyExclusive: prices: pd.Series = attr.ib(init=False) returns: pd.Series = attr.ib(init=False) trading_days_per_year: int = attr.ib(init=False, default=252) def __init__(self): # 直接调用构造函数就报错,引导用户用正确的方法 raise NotImplementedError("请使用from_prices或from_returns方法实例化对象") @classmethod def from_prices(cls, price_series: pd.Series, trading_days: int = 252): instance = cls.__new__(cls) # 创建实例但不触发__init__ instance.prices = price_series instance.returns = price_series.pct_change() instance.trading_days_per_year = trading_days return instance @classmethod def from_returns(cls, return_series: pd.Series, trading_days: int = 252): instance = cls.__new__(cls) instance.returns = return_series # 从基准值100开始计算对应的价格序列 instance.prices = pd.Series(100 * (1 + return_series).cumprod(), index=return_series.index) instance.trading_days_per_year = trading_days return instance if __name__ == "__main__": prices = pd.Series(data=[100, 101, 98, 104, 102, 108]) returns = pd.Series(data=[0.01, 0.03, -0.02, 0.01, -0.03, 0.04]) # 正常创建实例 obj_returns = MutuallyExclusive.from_returns(returns) obj_prices = MutuallyExclusive.from_prices(prices, trading_days=100) # 尝试直接调用常规构造函数会触发错误 try: obj = MutuallyExclusive(prices, returns) except NotImplementedError as e: print(e) # 输出:请使用from_prices或from_returns方法实例化对象
2. 保留常规构造函数但添加验证逻辑
如果你需要保留常规构造函数的灵活性(比如允许用户只传其中一个参数),可以在类初始化后添加验证逻辑,检查传入的prices和returns是否兼容,不兼容就抛出错误。
实现思路:
- 把
prices和returns的默认值设为None,允许只传其中一个; - 在
__attrs_post_init__方法中做三件事:检查是否同时传入了不兼容的参数、检查是否一个参数都没传、如果只传了一个则自动计算另一个。
代码示例:
import pandas as pd import attr from attrs import validators @attr.s class MutuallyExclusive: prices: pd.Series = attr.ib(default=None, validator=validators.optional(validators.instance_of(pd.Series))) returns: pd.Series = attr.ib(default=None, validator=validators.optional(validators.instance_of(pd.Series))) trading_days_per_year: int = attr.ib(default=252) @classmethod def from_prices(cls, price_series: pd.Series, trading_days: int = 252): return cls(price_series, price_series.pct_change(), trading_days) @classmethod def from_returns(cls, return_series: pd.Series, trading_days: int = 252): prices = pd.Series(100 * (1 + return_series).cumprod(), index=return_series.index) return cls(prices, return_series, trading_days) def __attrs_post_init__(self): # 检查是否同时传入了两个参数 if self.prices is not None and self.returns is not None: # 计算prices对应的returns,忽略第一个NaN值,对比精度到小数点后8位 calculated_returns = self.prices.pct_change() is_match = calculated_returns.iloc[1:].round(8).equals(self.returns.iloc[1:].round(8)) if not is_match: raise ValueError("传入的prices和returns不兼容,returns应为prices的pct_change结果") # 检查是否一个参数都没传 elif self.prices is None and self.returns is None: raise ValueError("必须传入prices或returns中的一个参数") # 只传了prices,自动计算returns elif self.prices is not None: self.returns = self.prices.pct_change() # 只传了returns,自动计算prices elif self.returns is not None: self.prices = pd.Series(100 * (1 + self.returns).cumprod(), index=self.returns.index) if __name__ == "__main__": prices = pd.Series(data=[100, 101, 98, 104, 102, 108]) returns = pd.Series(data=[0.01, 0.03, -0.02, 0.01, -0.03, 0.04]) # 正常实例化 obj_returns = MutuallyExclusive.from_returns(returns) obj_prices = MutuallyExclusive.from_prices(prices, trading_days=100) # 传入不兼容的参数会报错 try: obj = MutuallyExclusive(prices, returns) except ValueError as e: print(e) # 输出:传入的prices和returns不兼容... # 只传一个参数也能自动补全 obj_only_prices = MutuallyExclusive(prices=prices) obj_only_returns = MutuallyExclusive(returns=returns)
方案选择建议
- 如果你想完全避免用户误用常规构造函数,优先选第一种方案,彻底封死错误路径;
- 如果你需要给用户更多灵活性(比如允许只传一个参数直接实例化),第二种方案更合适,既能验证参数兼容性,又能自动补全缺失值。
内容的提问来源于stack exchange,提问作者Andi
相关产品推荐
相关产品推荐

