如何实现Pydantic模型的多态序列化?遇序列化反序列化问题求助
Pydantic子类序列化与反序列化问题修复
问题背景
尝试序列化包含基类NodeBase子类属性的Pydantic模型时,子类会被序列化为基类,且直接用基类验证子类字典会丢失类型信息,运行代码后出现序列化警告与断言错误。
原代码:
from typing import Dict, Literal, Union from pydantic import BaseModel, Field, RootModel class NodeBase(BaseModel): id: str class StartNode(NodeBase): type: Literal["start"] = "start" class EndNode(NodeBase): type: Literal["end"] = "end" class LLMNode(NodeBase): type: Literal["llm"] = "llm" name: str = Field(default_factory=lambda: id) purpose: str prompt: str model: Literal[ "gpt-4o", "gpt4-turbo", "gpt-4", "gpt-3.5-turbo", "azure-gpt-3.5-turbo" ] class NodeModel(RootModel): root: Union[StartNode, EndNode, LLMNode] class Graph(BaseModel): nodes: Dict[str, NodeModel] = Field(default_factory=dict) def add_node(self, node: Union[StartNode, EndNode, LLMNode]) -> None: self.nodes[node.id] = NodeModel(root=node) start_node = StartNode(id="start", type="start") llm_node = LLMNode(id="llm", type="llm", purpose="test", prompt="test", model="gpt-4o") end_node = EndNode(id="end", type="end") # ========= Node tests ========= start_node_dict = start_node.model_dump() llm_node_dict = llm_node.model_dump() end_node_dict = end_node.model_dump() # 是否可以用基类执行model_validate? start_node_from_dict = NodeBase.model_validate(start_node_dict) llm_node_from_dict = NodeBase.model_validate(llm_node_dict) end_node_from_dict = NodeBase.model_validate(end_node_dict) assert start_node == start_node_from_dict assert llm_node == llm_node_from_dict assert end_node == end_node_from_dict # ========= Graph tests ========= g = Graph() g.add_node(start_node) g.add_node(llm_node) g.add_node(end_node) g_dict = g.model_dump() g_from_dict = Graph.model_validate(g_dict) assert g == g_from_dict
错误信息
UserWarning: Pydantic序列化警告: 预期类型为`str`但得到`builtin_function_or_method` - 序列化结果可能不符合预期 return self.__pydantic_serializer__.to_python( Traceback (most recent call last): File "file.py", line 53, in <module> assert start_node == start_node_from_dict ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ AssertionError
需求
- 实现
Graph模型的正确序列化与反序列化,确保子类节点保留正确类型; - 支持通过
NodeBase.model_validate直接反序列化未知类型的NodeBase子类实例。
问题分析
- LLMNode字段错误:
name字段的默认工厂使用了Python内置函数id,生成的是函数对象而非字符串ID,触发序列化警告; - 类型丢失问题:直接用
NodeBase.model_validate处理子类字典时,Pydantic仅生成NodeBase实例,丢失子类属性与类型; - RootModel冗余:使用
RootModel包裹Union类型,增加了序列化层级,不利于类型自动解析。
修复方案
- 修复LLMNode的
name字段默认值,使用UUID生成合法字符串ID; - 给
NodeBase添加鉴别器配置,通过type字段自动识别子类; - 简化
Graph的节点字段类型,直接使用子类Union,利用鉴别器自动处理类型解析。
修复后完整代码
from typing import Dict, Literal, Union import uuid from pydantic import BaseModel, Field class NodeBase(BaseModel): id: str # 配置鉴别器,通过type字段自动匹配子类 model_config = {"discriminator": "type"} class StartNode(NodeBase): type: Literal["start"] = "start" class EndNode(NodeBase): type: Literal["end"] = "end" class LLMNode(NodeBase): type: Literal["llm"] = "llm" # 替换为UUID生成字符串ID,解决序列化警告 name: str = Field(default_factory=lambda: uuid.uuid4().hex) purpose: str prompt: str model: Literal[ "gpt-4o", "gpt4-turbo", "gpt-4", "gpt-3.5-turbo", "azure-gpt-3.5-turbo" ] # 定义节点类型Union,简化后续使用 NodeType = Union[StartNode, EndNode, LLMNode] class Graph(BaseModel): # 直接使用NodeType,鉴别器会自动处理子类的序列化/反序列化 nodes: Dict[str, NodeType] = Field(default_factory=dict) def add_node(self, node: NodeType) -> None: self.nodes[node.id] = node # ========= Node tests ========= start_node = StartNode(id="start") llm_node = LLMNode(id="llm", purpose="test", prompt="test", model="gpt-4o") end_node = EndNode(id="end") start_node_dict = start_node.model_dump() llm_node_dict = llm_node.model_dump() end_node_dict = end_node.model_dump() # 现在可通过NodeBase.model_validate直接得到对应子类实例 start_node_from_dict = NodeBase.model_validate(start_node_dict) llm_node_from_dict = NodeBase.model_validate(llm_node_dict) end_node_from_dict = NodeBase.model_validate(end_node_dict) # 验证类型与实例匹配 assert isinstance(start_node_from_dict, StartNode) assert start_node == start_node_from_dict assert isinstance(llm_node_from_dict, LLMNode) assert llm_node == llm_node_from_dict assert isinstance(end_node_from_dict, EndNode) assert end_node == end_node_from_dict # ========= Graph tests ========= g = Graph() g.add_node(start_node) g.add_node(llm_node) g.add_node(end_node) g_dict = g.model_dump() g_from_dict = Graph.model_validate(g_dict) # 验证反序列化后节点类型正确 assert isinstance(g_from_dict.nodes["start"], StartNode) assert isinstance(g_from_dict.nodes["llm"], LLMNode) assert isinstance(g_from_dict.nodes["end"], EndNode) assert g == g_from_dict
修复说明
- 鉴别器配置:
NodeBase的model_config中设置discriminator="type",Pydantic会根据type字段的值自动关联对应的子类,序列化时保留所有子类字段,反序列化时自动还原子类类型; - 字段默认值修复:用
uuid.uuid4().hex生成字符串ID,解决了序列化时的类型不匹配警告; - 基类验证支持:现在
NodeBase.model_validate可以接收任意子类的字典数据,自动返回对应的子类实例,满足未知类型节点的反序列化需求; - 简化结构:去掉冗余的
NodeModel,直接在Graph中使用子类Union类型,代码更简洁且类型解析更高效。
内容的提问来源于stack exchange,提问作者Asamaras
相关产品推荐
相关产品推荐

