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

如何在泛型类中使用Pydantic的from_orm方法?

问题描述

我正在编写一个基于Pydantic和SQLAlchemy的泛型仓储类,希望移除get方法中传入结果Pydantic模型作为参数的需求,现有代码如下:

class DatabaseRepository(Generic[T]):

    @classmethod
    async def get(cls, obj_id, model_class: Type[T]) -> T:
        table = cls.get_table()
        async with AsyncSession(cls.engine) as session:
            result = await session.get(table, obj_id)
        return model_class.from_orm(result)

我在网上了解到可以使用get_args获取泛型类传入的模型,但尝试后无法生效:

get_args(cls.__bases__)[0].from_orm(result)

cls.__bases__为空列表,无法访问到Pydantic模型,我也尝试过__orig_bases__,同样为空。注:T是继承自BaseModel的Pydantic模型。请问是否有办法移除上述model_class参数,仍能在泛型类中使用from_orm()方法?

解决方案

要在泛型仓储类中直接获取Pydantic模型类型T,可以通过以下两种实用方式实现:

方式一:利用泛型类的__orig_class__属性

Python 3.9+的泛型实现中,当实例化带具体类型参数的泛型子类时,__orig_class__会保留原始泛型类型信息,结合typing.get_args即可提取T:

from typing import Generic, TypeVar, Type, get_args
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel

T = TypeVar('T', bound=BaseModel)

class DatabaseRepository(Generic[T]):
    engine = ...  # 你的SQLAlchemy引擎实例
    table = ...   # 对应的SQLAlchemy表模型

    @classmethod
    async def get(cls, obj_id) -> T:
        # 提取泛型参数对应的具体Pydantic模型
        model_class = get_args(cls.__orig_class__)[0]
        async with AsyncSession(cls.engine) as session:
            result = await session.get(cls.table, obj_id)
        return model_class.from_orm(result)

# 使用示例
class UserModel(BaseModel):
    id: int
    name: str

class UserRepository(DatabaseRepository[UserModel]):
    table = User  # 关联的SQLAlchemy表类

# 调用时无需传入model_class参数
user = await UserRepository.get(1)

方式二:子类显式声明模型类型

如果担心__orig_class__的兼容性问题,可以在仓储子类中直接定义模型类型属性,父类直接读取该属性:

from typing import Generic, TypeVar, Type
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel

T = TypeVar('T', bound=BaseModel)

class DatabaseRepository(Generic[T]):
    engine = ...
    table = ...
    model_class: Type[T]  # 子类必须覆盖该属性

    @classmethod
    async def get(cls, obj_id) -> T:
        async with AsyncSession(cls.engine) as session:
            result = await session.get(cls.table, obj_id)
        return cls.model_class.from_orm(result)

# 使用示例
class UserRepository(DatabaseRepository[UserModel]):
    table = User
    model_class = UserModel

user = await UserRepository.get(1)

注意事项

  • 第一种方式依赖Python泛型运行时信息,在动态创建子类等复杂继承场景下可能失效。
  • 第二种方式更直观,兼容性更强,适合需要明确控制模型类型的场景。

内容的提问来源于stack exchange,提问作者Oren_C

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 05:21:30