Estoy usando Pydantic para definir datos jerárquicos en los que hay modelos con atributos idénticos.
Sin embargo, cuando guardo y cargo estos modelos, Pydantic ya no puede distinguir qué modelo se usó y elige el primero en la anotación de tipo de campo.
Entiendo que este es el comportamiento esperado según la documentación . Sin embargo, la información del tipo de clase es importante para mi aplicación.
¿Cuál es la forma recomendada de distinguir entre diferentes clases en Pydantic? Un truco es simplemente agregar un campo extraño a uno de los modelos, pero me gustaría encontrar una solución más elegante.
Vea el ejemplo simplificado a continuación: el container se inicializa con datos de tipo DataB , pero después de exportar y cargar, el nuevo container tiene datos de tipo DataA , ya que es el primer elemento en la declaración de tipo de container.data .
¡Gracias por tu ayuda!
from abc import ABC from pydantic import BaseModel #pydantic 1.8.2 from typing import Union class Data(BaseModel, ABC): """ base class for a Member """ number: float class DataA(Data): """ A type of Data""" pass class DataB(Data): """ Another type of Data """ pass class Container(BaseModel): """ container holds a subclass of Data """ data: Union[DataA, DataB] # initialize container with DataB data = DataB(number=1.0) container = Container(data=data) # export container to string and load new container from string string = container.json() new_container = Container.parse_raw(string) # look at type of container.data print(type(new_container.data).__name__) # >>> DataAComo se señaló correctamente en los comentarios, sin almacenar información adicional, los modelos no se pueden distinguir al analizar.
A partir de hoy (pydantic v1.8.2), la forma más canónica de distinguir modelos al analizar en una Union (en caso de ambigüedad) es agregar explícitamente un especificador de tipo Literal . Se verá así:
from abc import ABC from pydantic import BaseModel from typing import Union, Literal class Data(BaseModel, ABC): """ base class for a Member """ number: float class DataA(Data): """ A type of Data""" tag: Literal['A'] = 'A' class DataB(Data): """ Another type of Data """ tag: Literal['B'] = 'B' class Container(BaseModel): """ container holds a subclass of Data """ data: Union[DataA, DataB] # initialize container with DataB data = DataB(number=1.0) container = Container(data=data) # export container to string and load new container from string string = container.json() new_container = Container.parse_raw(string) # look at type of container.data print(type(new_container.data).__name__) # >>> DataBEste método se puede automatizar, pero puedes usarlo bajo tu propia responsabilidad, ya que rompe el tipeo estático y usa objetos que pueden cambiar en futuras versiones:
from pydantic.fields import ModelField class Data(BaseModel, ABC): """ base class for a Member """ number: float def __init_subclass__(cls, **kwargs): name = 'tag' value = cls.__name__ annotation = Literal[value] tag_field = ModelField.infer(name=name, value=value, annotation=annotation, class_validators=None, config=cls.__config__) cls.__fields__[name] = tag_field cls.__annotations__[name] = annotation class DataA(Data): """ A type of Data""" pass class DataB(Data): """ Another type of Data """ passMientras tanto, estoy tratando de piratear algo juntos usando validadores personalizados. Básicamente, el decorador de clases agrega un campo class_name: str , que se agrega a la cadena json. Luego, el validador busca la subclase correcta en función de su valor.
def register_distinct_subclasses(fields: tuple): """ fields is tuple of subclasses that we want to be registered as distinct """ field_map = {field.__name__: field for field in fields} def _register_distinct_subclasses(cls): """ cls is the superclass of fields, which we add a new validator to """ orig_init = cls.__init__ class _class: class_name: str def __init__(self, **kwargs): class_name = type(self).__name__ kwargs["class_name"] = class_name orig_init(**kwargs) @classmethod def __get_validators__(cls): yield cls.validate @classmethod def validate(cls, v): if isinstance(v, dict): class_name = v.get("class_name") json_string = json.dumps(v) else: class_name = v.class_name json_string = v.json() cls_type = field_map[class_name] return cls_type.parse_raw(json_string) return _class return _register_distinct_subclassesque se llama de la siguiente manera
Data = register_distinct_subclasses((DataA, DataB))(Data)