Refactor field serialization and add model

serialization in Template and FrontendNode classes
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-12-09 18:11:51 -03:00
commit 3c955de5e3
3 changed files with 69 additions and 48 deletions

View file

@ -1,11 +1,11 @@
from abc import ABC from abc import ABC
from typing import Any, Optional, Union from typing import Any, Optional, Union
from pydantic import BaseModel from pydantic import BaseModel, Field, field_serializer, model_serializer
class TemplateFieldCreator(BaseModel, ABC): class TemplateFieldCreator(BaseModel, ABC):
field_type: str = "str" field_type: str = Field(default="str", alias="type")
"""The type of field this is. Default is a string.""" """The type of field this is. Default is a string."""
required: bool = False required: bool = False
@ -14,7 +14,7 @@ class TemplateFieldCreator(BaseModel, ABC):
placeholder: str = "" placeholder: str = ""
"""A placeholder string for the field. Default is an empty string.""" """A placeholder string for the field. Default is an empty string."""
is_list: bool = False is_list: bool = Field(default=False, alias="list")
"""Defines if the field is a list. Default is False.""" """Defines if the field is a list. Default is False."""
show: bool = True show: bool = True
@ -26,7 +26,7 @@ class TemplateFieldCreator(BaseModel, ABC):
value: Any = None value: Any = None
"""The value of the field. Default is None.""" """The value of the field. Default is None."""
file_types: list[str] = [] file_types: list[str] = Field(default=[], alias="fileTypes")
"""List of file types associated with the field. Default is an empty list. (duplicate)""" """List of file types associated with the field. Default is an empty list. (duplicate)"""
file_path: Union[str, None] = None file_path: Union[str, None] = None
@ -35,7 +35,7 @@ class TemplateFieldCreator(BaseModel, ABC):
password: bool = False password: bool = False
"""Specifies if the field is a password. Defaults to False.""" """Specifies if the field is a password. Defaults to False."""
options: list[str] = [] options: list[str] = None
"""List of options for the field. Only used when is_list=True. Default is an empty list.""" """List of options for the field. Only used when is_list=True. Default is an empty list."""
name: str = "" name: str = ""
@ -47,7 +47,7 @@ class TemplateFieldCreator(BaseModel, ABC):
advanced: bool = False advanced: bool = False
"""Specifies if the field will an advanced parameter (hidden). Defaults to False.""" """Specifies if the field will an advanced parameter (hidden). Defaults to False."""
input_types: list[str] = [] input_types: Optional[list[str]] = None
"""List of input types for the handle when the field has more than one type. Default is an empty list.""" """List of input types for the handle when the field has more than one type. Default is an empty list."""
dynamic: bool = False dynamic: bool = False
@ -59,22 +59,31 @@ class TemplateFieldCreator(BaseModel, ABC):
refresh: Optional[bool] = None refresh: Optional[bool] = None
"""Specifies if the field should be refreshed. Defaults to False.""" """Specifies if the field should be refreshed. Defaults to False."""
def to_dict(self): @model_serializer(mode="wrap")
result = self.model_dump() def serialize_model(self, handler):
# Remove key if it is None # This will be the result of model_dump or dict()
for key in list(result.keys()): # so we need to build a dict to return
if result[key] is None or result[key] == [] and key != "value": result = handler(self)
del result[key] result["value"] = self.value
result["type"] = result.pop("field_type")
result["list"] = result.pop("is_list")
if result.get("file_types"):
result["fileTypes"] = result.pop("file_types")
if self.field_type == "file":
result["file_path"] = self.file_path
return result return result
# for key in list(result.keys()):
# if result[key] is None or result[key] == [] and key != "value":
# del result[key]
# return result
def to_dict(self):
return self.model_dump(by_alias=True, exclude_none=True)
@field_serializer("file_path")
def serialize_file_path(self, value):
if self.field_type == "file":
return value
return None
class TemplateField(TemplateFieldCreator): class TemplateField(TemplateFieldCreator):
pass pass

View file

@ -3,11 +3,12 @@ from collections import defaultdict
from typing import ClassVar, Dict, List, Optional from typing import ClassVar, Dict, List, Optional
from langflow.template.field.base import TemplateField from langflow.template.field.base import TemplateField
from langflow.template.frontend_node.constants import CLASSES_TO_REMOVE, FORCE_SHOW_FIELDS from langflow.template.frontend_node.constants import (CLASSES_TO_REMOVE,
FORCE_SHOW_FIELDS)
from langflow.template.frontend_node.formatter import field_formatters from langflow.template.frontend_node.formatter import field_formatters
from langflow.template.template.base import Template from langflow.template.template.base import Template
from langflow.utils import constants from langflow.utils import constants
from pydantic import BaseModel, Field from pydantic import BaseModel, Field, field_serializer, model_serializer
class FieldFormatters(BaseModel): class FieldFormatters(BaseModel):
@ -63,26 +64,31 @@ class FrontendNode(BaseModel):
"""Sets the documentation of the frontend node.""" """Sets the documentation of the frontend node."""
self.documentation = documentation self.documentation = documentation
def process_base_classes(self) -> None: @field_serializer("base_classes")
def process_base_classes(self, base_classes: List[str]) -> List[str]:
"""Removes unwanted base classes from the list of base classes.""" """Removes unwanted base classes from the list of base classes."""
self.base_classes = [base_class for base_class in self.base_classes if base_class not in CLASSES_TO_REMOVE]
return [base_class for base_class in base_classes if base_class not in CLASSES_TO_REMOVE]
@field_serializer("display_name")
def process_display_name(self, display_name: str) -> str:
"""Sets the display name of the frontend node."""
return display_name or self.name
@model_serializer(mode="wrap")
def serialize(self, handler):
result = handler(self)
result["template"] = self.template.to_dict(self.format_field)
name = result.pop("name")
return {name: result}
def to_dict(self) -> dict: def to_dict(self) -> dict:
"""Returns a dict representation of the frontend node.""" """Returns a dict representation of the frontend node."""
self.process_base_classes()
return { return self.model_dump(by_alias=True, exclude_none=True)
self.name: {
"template": self.template.to_dict(self.format_field),
"description": self.description,
"base_classes": self.base_classes,
"display_name": self.display_name or self.name,
"custom_fields": self.custom_fields,
"output_types": self.output_types,
"documentation": self.documentation,
"beta": self.beta,
"error": self.error,
},
}
def add_extra_fields(self) -> None: def add_extra_fields(self) -> None:
pass pass

View file

@ -1,9 +1,8 @@
from typing import Callable, Optional, Union from typing import Callable, Union
from pydantic import BaseModel
from langflow.template.field.base import TemplateField from langflow.template.field.base import TemplateField
from langflow.utils.constants import DIRECT_TYPES from langflow.utils.constants import DIRECT_TYPES
from pydantic import BaseModel, model_serializer
class Template(BaseModel): class Template(BaseModel):
@ -12,12 +11,11 @@ class Template(BaseModel):
def process_fields( def process_fields(
self, self,
name: Optional[str] = None,
format_field_func: Union[Callable, None] = None, format_field_func: Union[Callable, None] = None,
): ):
if format_field_func: if format_field_func:
for field in self.fields: for field in self.fields:
format_field_func(field, name) format_field_func(field, self.type_name)
def sort_fields(self): def sort_fields(self):
# first sort alphabetically # first sort alphabetically
@ -25,12 +23,20 @@ class Template(BaseModel):
self.fields.sort(key=lambda x: x.name) self.fields.sort(key=lambda x: x.name)
self.fields.sort(key=lambda x: x.field_type in DIRECT_TYPES, reverse=False) self.fields.sort(key=lambda x: x.field_type in DIRECT_TYPES, reverse=False)
def to_dict(self, format_field_func=None): @model_serializer(mode="wrap")
self.process_fields(self.type_name, format_field_func) def serialize_model(self, handler):
self.sort_fields() result = handler(self)
result = {field.name: field.to_dict() for field in self.fields} for field in self.fields:
result["_type"] = self.type_name # type: ignore result[field.name] = field.to_dict()
result["_type"] = result.pop("type_name")
return result return result
def to_dict(self, format_field_func=None):
self.process_fields(format_field_func)
self.sort_fields()
# result = {field.name: field.to_dict() for field in self.fields}
# result["_type"] = self.type_name # type: ignore
return self.model_dump(by_alias=True, exclude_none=True, exclude={"fields"})
def add_field(self, field: TemplateField) -> None: def add_field(self, field: TemplateField) -> None:
self.fields.append(field) self.fields.append(field)