如何为配对层级的数据文件类添加类型变量注解以适配子类?
为带继承层级的DataFile/Header结构添加类型注解
需求概述
- 旧代码库采用
File(Header, Data)结构,DataFile与Header存在配对继承层级(新版本子类继承旧版本) - 要求类型检查器能正确识别:
DataFileV1().header为HeaderV1类型,DataFileV2().header为HeaderV2类型 DataFileV1必须可直接实例化,且作为DataFileV2的父类- 希望避免过多样板代码,同时了解从头设计这类优先考虑类型注解的API的思路
原代码示例
class MetaDataMixin: def __init__(self, metadata=None, *args, **kwargs): super().__init__(*args, **kwargs) self.metadata = {} if metadata: self.metadata.update(metadata) class DataFile: def __init__(self, header=None, data=None): self.header = header or self.header_class() class HeaderV1: magic_number = b'HDR' format_version = 1 class DataFileV1(DataFile): header_class = HeaderV1 class HeaderV2(HeaderV1, MetaDataMixin): format_version = 2 class DataFileV2(DataFileV1): header_class = HeaderV2
尝试的代码及问题
尝试用泛型类绑定Header类型,但子类DataFileV2继承DataFileV1后,类型检查器无法识别header已升级为HeaderV2:
import typing as ty class Header: pass HdrT = ty.TypeVar('HdrT', bound=Header) class MetaDataMixin: metadata: dict[str, str] def __init__(self, metadata=None, *args, **kwargs): super().__init__(*args, **kwargs) self.metadata = {} if metadata: self.metadata.update(metadata) class DataFile(ty.Generic[HdrT]): header: HdrT header_class: type[HdrT] def __init__(self, header: HdrT | None = None, data: ty.Any = None): self.header = header or self.header_class() class HeaderV1(Header): magic_number: bytes = b'HDR' format_version: int = 1 class DataFileV1(DataFile[HeaderV1]): header_class = HeaderV1 class HeaderV2(HeaderV1, MetaDataMixin): format_version = 2 class DataFileV2(DataFileV1, DataFile[HeaderV2]): header_class = HeaderV2 file1 = DataFileV1() file2 = DataFileV2() print(file2.header.format_version) if ty.TYPE_CHECKING: # 类型检查器显示为HeaderV1,但期望是HeaderV2 reveal_type(file2.header) else: # 运行时实际是HeaderV2 print(file2.header.metadata)
解决方案
改造现有代码:调整泛型继承顺序
问题出在多重继承的MRO顺序,类型检查器优先使用排在后面的泛型父类注解。将DataFile[HeaderV2]放在DataFileV1前面,让类型检查器优先识别新的Header类型:
import typing as ty HdrT = ty.TypeVar('HdrT', bound='Header') class Header: pass class MetaDataMixin: metadata: dict[str, str] def __init__(self, metadata: dict[str, str] | None = None, *args, **kwargs): super().__init__(*args, **kwargs) self.metadata = metadata.copy() if metadata else {} class DataFile(ty.Generic[HdrT]): header_class: type[HdrT] header: HdrT def __init__(self, header: HdrT | None = None, data: ty.Any = None): self.header = header or self.header_class() class HeaderV1(Header): magic_number: bytes = b'HDR' format_version: int = 1 class DataFileV1(DataFile[HeaderV1]): header_class = HeaderV1 class HeaderV2(HeaderV1, MetaDataMixin): format_version: int = 2 # 调整继承顺序,让DataFile[HeaderV2]排在前面 class DataFileV2(DataFile[HeaderV2], DataFileV1): header_class = HeaderV2 file1 = DataFileV1() file2 = DataFileV2() if ty.TYPE_CHECKING: reveal_type(file1.header) # 正确识别为HeaderV1 reveal_type(file2.header) # 正确识别为HeaderV2
从头设计类型友好的API
如果可以重新设计API,推荐以下思路减少样板代码并强化类型安全:
使用泛型基类+
__init_subclass__自动绑定Header类型
用__init_subclass__简化子类定义,避免手动重复指定泛型参数:import typing as ty HdrT = ty.TypeVar('HdrT', bound='BaseHeader') class BaseHeader: format_version: int class MetaDataMixin: metadata: dict[str, str] def __init__(self, metadata: dict[str, str] | None = None, *args, **kwargs): super().__init__(*args, **kwargs) self.metadata = metadata.copy() if metadata else {} class BaseDataFile(ty.Generic[HdrT]): header_type: type[HdrT] def __init_subclass__(cls, header_type: type[HdrT], **kwargs): super().__init_subclass__(**kwargs) cls.header_type = header_type def __init__(self, header: HdrT | None = None, data: ty.Any = None): self.header = header or self.header_type() self.data = data or [] # 子类定义简洁,无需手动指定泛型参数 class HeaderV1(BaseHeader): magic_number: bytes = b'HDR' format_version: int = 1 class DataFileV1(BaseDataFile, header_type=HeaderV1): pass class HeaderV2(HeaderV1, MetaDataMixin): format_version: int = 2 # 继承DataFileV1的同时绑定新的Header类型 class DataFileV2(DataFileV1, BaseDataFile, header_type=HeaderV2): pass结合数据类减少样板代码
使用dataclasses自动生成构造函数和类型注解,让代码更简洁:import typing as ty from dataclasses import dataclass, field HdrT = ty.TypeVar('HdrT', bound='BaseHeader') @dataclass class BaseHeader: format_version: int @dataclass class HeaderV1(BaseHeader): magic_number: bytes = b'HDR' format_version: int = 1 @dataclass class HeaderV2(HeaderV1): metadata: dict[str, str] = field(default_factory=dict) format_version: int = 2 class BaseDataFile(ty.Generic[HdrT]): def __init__(self, header: HdrT | None = None, data: list[ty.Any] = None): self.header = header or self._create_default_header() self.data = data or [] def _create_default_header(self) -> HdrT: raise NotImplementedError class DataFileV1(BaseDataFile[HeaderV1]): def _create_default_header(self) -> HeaderV1: return HeaderV1() class DataFileV2(BaseDataFile[HeaderV2], DataFileV1): def _create_default_header(self) -> HeaderV2: return HeaderV2()
内容的提问来源于stack exchange,提问作者Chris Markiewicz
相关产品推荐
相关产品推荐

