如何在Python中实现基于子类泛型类型的验证基方法
可以将
get_all泛型化到基类,彻底消除子类重复代码 当然可以。借助Python泛型的类型绑定特性,你完全可以把get_all的通用逻辑放到BaseDataStore基类中,子类只需指定对应的Pydantic模型即可,无需重复实现该方法。下面分两种场景给出具体实现:
场景1:所有子类的客户端数据获取逻辑完全一致
如果所有子类调用客户端的方式相同(比如都用client.get_all()),可以直接在基类中完成全部通用逻辑:
from abc import ABC from typing import Generic, TypeVar, List, get_args from pydantic import BaseModel # 基础Pydantic模型 class CustomBaseModel(BaseModel): # 定义通用字段 pass # 绑定CustomBaseModel的TypeVar BoundedModel = TypeVar('BoundedModel', bound=CustomBaseModel) # 具体Pydantic模型子类 class Metadata(CustomBaseModel): id: int name: str class Transcript(CustomBaseModel): transcript_id: str content: str # 泛型基类 class BaseDataStore(ABC, Generic[BoundedModel]): def __init__(self, client): self.client = client # 自动获取子类绑定的具体Pydantic模型 self.model_class = get_args(self.__orig_class__)[0] def get_all(self) -> List[BoundedModel]: # 调用客户端获取原始数据 raw_data = self.client.get_all() # 用绑定的模型验证并返回列表 return [self.model_class.parse_obj(item) for item in raw_data] # 子类仅需指定泛型参数,无需实现get_all class MetadataStore(BaseDataStore[Metadata]): pass class TranscriptStore(BaseDataStore[Transcript]): pass
核心原理:get_args(self.__orig_class__)[0]会自动读取子类继承时指定的泛型参数(比如MetadataStore对应的Metadata),让基类动态知道该用哪个模型解析数据。
场景2:子类的客户端数据获取逻辑不同
如果不同子类调用客户端的方法不同(比如一个用client.fetch_metadata(),一个用client.get_transcripts()),可以把数据获取逻辑抽成抽象方法,基类只负责通用的模型解析:
from abc import ABC, abstractmethod from typing import Generic, TypeVar, List, get_args, Any from pydantic import BaseModel # 重复的基础定义(CustomBaseModel、BoundedModel、Metadata、Transcript同上) class BaseDataStore(ABC, Generic[BoundedModel]): def __init__(self, client): self.client = client self.model_class = get_args(self.__orig_class__)[0] def get_all(self) -> List[BoundedModel]: # 调用子类实现的抽象方法获取原始数据 raw_data = self._fetch_raw_data() return [self.model_class.parse_obj(item) for item in raw_data] @abstractmethod def _fetch_raw_data(self) -> List[Any]: # 子类需实现具体的数据获取逻辑 pass # 子类仅需实现数据获取的抽象方法 class MetadataStore(BaseDataStore[Metadata]): def _fetch_raw_data(self) -> List[Any]: return self.client.fetch_metadata() class TranscriptStore(BaseDataStore[Transcript]): def _fetch_raw_data(self) -> List[Any]: return self.client.get_transcripts()
这种方式既保留了get_all的通用验证逻辑,又允许子类自定义数据获取细节,避免了重复编写模型解析代码。
注意事项
- 确保Python版本在3.8及以上,
get_args和__orig_class__在该版本及之后才能正常工作。 - 如果客户端是类级别的静态属性,可调整
model_class的获取逻辑,比如通过get_args(self.__orig_bases__[0])在类初始化时读取。
内容的提问来源于stack exchange,提问作者alexcs
相关产品推荐
相关产品推荐

