如何使抽象类强制其实现类必须为dataclass类型
如何强制继承CsvableDataclass的类必须是@dataclass?
我定义了一个抽象类CsvableDataclass,用来将实现该类的类型实例列表转换为CSV字符串,代码如下:
from abc import ABC, abstractmethod from dataclasses import dataclass, fields from typing import List, Type, TypeVar T = TypeVar('T') class CsvableDataclass(ABC): @classmethod def to_csv_header(cls) -> str: return ','.join(f.name for f in fields(cls)) def to_csv_row(self) -> str: return ','.join(self.format_field(f.name) for f in fields(self)) @abstractmethod def format_field(self, field_name: str) -> str: pass # 生成完整CSV字符串 @staticmethod def to_csv_str(t: Type[T], data: List[T]): return '\n'.join([t.to_csv_header()] + [r.to_csv_row() for r in data]) @dataclass class A(CsvableDataclass): x: int y: int def format_field(self, field_name: str) -> str: if field_name == 'x': return str(self.x) if field_name == 'y': return str(self.y) raise Exception(f"invalid field_name: {field_name}") CsvableDataclass.to_csv_str(A, [A(1,2),A(3,4)]) # 输出结果:"x,y\n1,2\n3,4"
问题在于dataclasses.fields()只对被@dataclass装饰的类/实例生效,我希望通过类型注解强制所有继承CsvableDataclass的子类必须是dataclass,该怎么做?
可行方案
1. 用Protocol+TypeVar做静态类型约束
我们可以定义一个协议(Protocol)来标识dataclass的特征——所有dataclass都会自动生成__dataclass_fields__属性,利用这一点来约束子类:
from abc import ABC, abstractmethod from dataclasses import dataclass, fields, is_dataclass from typing import List, Type, TypeVar, Protocol # 定义协议,标记具备dataclass特征的类 class DataclassProtocol(Protocol): __dataclass_fields__: dict # 重新定义TypeVar,绑定到同时继承CsvableDataclass和DataclassProtocol的类型 T = TypeVar('T', bound='CsvableDataclass') class CsvableDataclass(ABC, DataclassProtocol): @classmethod def to_csv_header(cls) -> str: # 额外加运行时检查,防止静态检查被绕过 if not is_dataclass(cls): raise TypeError(f"{cls.__name__} 必须是@dataclass装饰的类") return ','.join(f.name for f in fields(cls)) def to_csv_row(self) -> str: return ','.join(self.format_field(f.name) for f in fields(self)) @abstractmethod def format_field(self, field_name: str) -> str: pass @staticmethod def to_csv_str(t: Type[T], data: List[T]): return '\n'.join([t.to_csv_header()] + [r.to_csv_row() for r in data])
这样一来,只要子类没加@dataclass,mypy这类静态类型检查器就会报错,提示子类缺少__dataclass_fields__属性,不符合协议要求。
2. 用__init_subclass__做运行时强制约束
如果想同时在运行时也确保子类是dataclass,可以在父类的__init_subclass__方法里加检查,子类继承时就会触发:
from abc import ABC, abstractmethod from dataclasses import dataclass, fields, is_dataclass from typing import List, Type, TypeVar T = TypeVar('T') class CsvableDataclass(ABC): def __init_subclass__(cls) -> None: super().__init_subclass__() if not is_dataclass(cls): raise TypeError(f"子类 {cls.__name__} 必须使用@dataclass装饰") @classmethod def to_csv_header(cls) -> str: return ','.join(f.name for f in fields(cls)) def to_csv_row(self) -> str: return ','.join(self.format_field(f.name) for f in fields(self)) @abstractmethod def format_field(self, field_name: str) -> str: pass @staticmethod def to_csv_str(t: Type[T], data: List[T]): return '\n'.join([t.to_csv_header()] + [r.to_csv_row() for r in data])
这种方式不管静态检查过没过,只要子类没加@dataclass,程序运行时一加载子类就会抛出错误,彻底避免后续调用fields()时崩溃。
3. 简化版:用mypy的predicate做类型检查
如果你只用mypy做静态检查,可以用mypy_extensions的@predicate装饰器来直接约束:
from abc import ABC, abstractmethod from dataclasses import dataclass, fields, is_dataclass from typing import List, Type, TypeVar from mypy_extensions import predicate @predicate def is_dataclass_type(cls) -> bool: return is_dataclass(cls) T = TypeVar('T', bound='CsvableDataclass') class CsvableDataclass(ABC): @classmethod def to_csv_header(cls: Type[T]) -> str: assert is_dataclass_type(cls), f"{cls.__name__} must be a dataclass" return ','.join(f.name for f in fields(cls)) # 其余方法保持不变
不过这个方案需要安装mypy_extensions,而且只对mypy生效,没有运行时检查。
内容的提问来源于stack exchange,提问作者Captain_Obvious
相关产品推荐
相关产品推荐

