You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何从泛型传入类型自动设置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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 12:25:08