如何从泛型传入类型自动设置BaseRepository的model_class属性?
自动从泛型参数提取model_class的实现方案
你可以通过解析子类继承时的泛型基类元信息,在__init_subclass__中自动提取泛型参数并赋值给model_class,无需手动指定。具体实现如下:
完整代码示例
from typing import Generic, TypeVar, get_args, get_origin from app.core.models.tables import Base, SomeModel BaseSQLAlchemyModel = TypeVar('BaseSQLAlchemyModel', bound=Base) class BaseRepository(Generic[BaseSQLAlchemyModel]): model_class: type[BaseSQLAlchemyModel] def __init_subclass__(cls) -> None: # 遍历子类的原始基类(保留泛型参数的版本) for base in cls.__orig_bases__: # 确认当前基类是BaseRepository的泛型实例 if get_origin(base) is BaseRepository: # 提取泛型参数列表 generic_args = get_args(base) if generic_args: model = generic_args[0] # 验证参数是Base的子类 if issubclass(model, Base): cls.model_class = model return # 如果没找到有效泛型参数,抛出异常 raise ValueError("子类必须继承带具体SQLAlchemy模型的BaseRepository泛型") def some_base_method(self) -> BaseSQLAlchemyModel: # 示例方法:使用model_class操作数据库 return self.model_class() # 这里替换为实际的SQLAlchemy操作 # 子类无需手动设置model_class class MyRepository(BaseRepository[SomeModel]): pass
关键逻辑说明
__orig_bases__:Python为泛型子类自动生成的属性,存储了子类继承时的原始基类(包含具体泛型参数的版本)。get_origin(base):获取泛型类型的原始类(比如get_origin(BaseRepository[SomeModel])会返回BaseRepository)。get_args(base):提取泛型参数列表(比如get_args(BaseRepository[SomeModel])会返回(SomeModel,))。- 类型验证:确保提取到的泛型参数是
Base的子类,避免传入非法类型。
这样,子类在继承BaseRepository[XXXModel]时,model_class会被自动赋值为XXXModel,无需手动声明。
内容的提问来源于stack exchange,提问作者Dmitriy Lunev
相关产品推荐
相关产品推荐

