Python泛型仓库类中如何针对特定泛型模型实现属性检查?
Python泛型仓库类中如何针对特定泛型模型实现属性检查?
你遇到的其实是Python泛型的一个常见小坑:在类定义阶段,ModelType还只是个TypeVar标记,并不是你实例化时传入的User类。所以BaseRepository里写的model = ModelType根本没把model绑定到具体的User类型上,调用list方法时检查self.model的属性自然会失败。
下面给你几种可行的解决方案,按需选择:
方法一:为每个模型创建专属仓库子类
这是最直观易维护的方式,给每个模型写对应的仓库子类,在子类里明确指定model属性:
# SQLAlchemy models from sqlalchemy.orm import DeclarativeBase from typing import TypeVar, Generic class BaseModel(DeclarativeBase): pass class User(BaseModel): name: str ModelType = TypeVar("ModelType", bound=BaseModel) class BaseRepository(Generic[ModelType]): # 这里只做类型注解,具体赋值交给子类 model: type[ModelType] def list(self, name_filter: str): if not hasattr(self.model, "name"): raise Exception(f"模型 {self.model.__name__} 没有 'name' 属性") # 这里可以继续编写查询逻辑,比如: # return session.query(self.model).filter(self.model.name == name_filter).all() # 给User模型创建专属仓库 class UserRepository(BaseRepository[User]): model = User # 使用起来也很简单 repo = UserRepository() repo.list("John") # 现在不会再报错啦
这种方式的好处是结构清晰,后续要给特定模型扩展仓库方法也很方便。
方法二:实例化时自动获取泛型参数
如果不想为每个模型都写子类,可以利用Python的typing.get_args,在__init__里动态拿到当前实例绑定的具体泛型类型:
# SQLAlchemy models from sqlalchemy.orm import DeclarativeBase from typing import TypeVar, Generic, get_args class BaseModel(DeclarativeBase): pass class User(BaseModel): name: str ModelType = TypeVar("ModelType", bound=BaseModel) class BaseRepository(Generic[ModelType]): def __init__(self): # 获取泛型参数 generic_args = get_args(self.__orig_class__) if not generic_args: raise ValueError("BaseRepository 必须指定泛型类型参数") self.model = generic_args[0] def list(self, name_filter: str): if not hasattr(self.model, "name"): raise Exception(f"模型 {self.model.__name__} 没有 'name' 属性") # 后续查询逻辑... # 直接实例化时指定泛型即可 repo = BaseRepository[User]() repo.list("John") # 正常工作
这里的self.__orig_class__是Python为泛型实例保留的原始类信息,通过get_args就能取出你传入的User类型,实现动态绑定。
额外优化:静态检查提前发现问题
如果想在写代码的时候就发现模型缺少name属性的问题,可以定义一个协议(Protocol),让ModelType绑定到这个协议上:
from typing import Protocol, TypeVar, Generic # 定义一个要求有name属性的协议 class HasName(Protocol): name: str # 让ModelType同时绑定BaseModel和HasName ModelType = TypeVar("ModelType", bound=BaseModel & HasName) class BaseRepository(Generic[ModelType]): # ... 后续代码不变
这样如果传入的模型没有name属性,IDE或者mypy这类静态检查工具会直接提示错误,不用等到运行时才踩坑。
备注:内容来源于stack exchange,提问作者Kamil Saitov
相关产品推荐
相关产品推荐

