如何自动注册Pydantic可识别的Discriminated Union(区分联合类型)
动态配置Pydantic区分联合类型(Discriminated Union)
问题背景
我使用Pydantic 1.10.7和Python 3.11.2,需要实现一个递归结构的Pydantic模型,通过区分联合类型自动完成子类的反序列化。最初的实现需要手动在Union中列出所有子类,不仅存在类型提示循环依赖问题,新增子类时还容易遗漏添加到联合类型里。
最初的手动实现代码:
from pydantic import BaseModel, Field from typing import Annotated, List, Union, Literal class Base(BaseModel): kind: str sub_models: Annotated[ List[Union[A,B]], Field( default_factory=list, discriminator="kind" ) ] class A(Base): kind: Literal["a"] a_field: str class B(Base): kind: Literal["b"] b_field: str
解决方案
通过__init_subclass__钩子自动注册子类,并动态更新父类的联合类型注解,同时处理循环引用问题:
from pydantic import BaseModel, Field from typing import Annotated, List, Union, Literal, TypeVar, ForwardRef, Set, Type # 定义TypeVar和ForwardRef处理循环引用 B = TypeVar("B", bound="Base") BaseRef = ForwardRef("Base") class Base(BaseModel): kind: str sub_models: Annotated[ List[Union[BaseRef]], Field(default_factory=list, discriminator="kind") ] _subs: Set[Type[B]] = set() def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) # 注册子类 Base._subs.add(cls) # 自动设置子类的kind字段为类名小写的Literal类型 cls.__annotations__["kind"] = Literal[cls.__name__.lower()] # 动态生成包含所有子类的Union类型 union_type = Union[tuple(Base._subs)] # 更新Base的sub_models注解 Base.__annotations__["sub_models"] = Annotated[ List[union_type], Field(default_factory=list, discriminator="kind") ] # 重新构建模型的字段,让Pydantic识别更新后的注解 Base.model_rebuild() class A(Base): a_field: str class B(Base): b_field: str # 最后再重建一次Base模型,确保所有子类都被注册 Base.model_rebuild()
关键说明
- 循环引用处理:使用
ForwardRef延迟解析Base类型,避免定义时的循环依赖错误。 - 自动注册子类:
__init_subclass__会在每个子类定义时触发,自动将子类添加到_subs集合中。 - 动态更新联合类型:每次新增子类后,重新生成包含所有已注册子类的
Union类型,并更新sub_models的注解,再调用model_rebuild()让Pydantic重新解析模型结构。 - 自动设置kind字段:子类的
kind会自动设为类名小写的Literal类型,无需手动编写。
测试验证
# 测试反序列化 test_data = { "kind": "a", "a_field": "test_a", "sub_models": [ { "kind": "b", "b_field": "test_b", "sub_models": [] } ] } instance = Base.parse_obj(test_data) print(instance) # 输出: kind='a' a_field='test_a' sub_models=[kind='b' b_field='test_b' sub_models=[]] print(type(instance)) # <class '__main__.A'> print(type(instance.sub_models[0])) # <class '__main__.B'>
内容的提问来源于stack exchange,提问作者abstrus
相关产品推荐
相关产品推荐

