🔨 refactor(base.py): refactor FrontendNode.format_field() method to improve readability and maintainability (#363)

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-05-27 14:10:38 -03:00 • committed by GitHub
commit c224608601
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 138 additions and 134 deletions

View file

@ -42,7 +42,7 @@ class LangChainTypeCreator(BaseModel, ABC):
# so we should update the result dict # so we should update the result dict
node = self.frontend_node(name) node = self.frontend_node(name)
if node is not None: if node is not None:
node = node.to_dict() node = node.to_dict() # type: ignore
result[self.type_name].update(node) result[self.type_name].update(node)
return result return result

View file

@ -1,5 +1,6 @@
import re
from abc import ABC from abc import ABC
from typing import Any, Callable, Dict, Optional, Union from typing import Any, Callable, List, Optional, Union
from pydantic import BaseModel from pydantic import BaseModel
@ -41,76 +42,6 @@ class TemplateFieldCreator(BaseModel, ABC):
result["content"] = self.content result["content"] = self.content
return result return result
def process_field(
self, key: str, value: Dict[str, Any], name: Optional[str] = None
) -> None:
_type = value["type"]
# Remove 'Optional' wrapper
if "Optional" in _type:
_type = _type.replace("Optional[", "")[:-1]
# Check for list type
if "List" in _type:
_type = _type.replace("List[", "")[:-1]
self.is_list = True
# Replace 'Mapping' with 'dict'
if "Mapping" in _type:
_type = _type.replace("Mapping", "dict")
# Change type from str to Tool
self.field_type = "Tool" if key in {"allowed_tools"} else self.field_type
self.field_type = "int" if key in {"max_value_length"} else self.field_type
# Show or not field
self.show = bool(
(self.required and key not in ["input_variables"])
or key in FORCE_SHOW_FIELDS
or "api_key" in key
)
# Add password field
self.password = any(
text in key.lower() for text in {"password", "token", "api", "key"}
)
# Add multline
self.multiline = key in {
"suffix",
"prefix",
"template",
"examples",
"code",
"headers",
}
# Replace dict type with str
if "dict" in self.field_type.lower():
self.field_type = "code"
if key == "dict_":
self.field_type = "file"
self.suffixes = [".json", ".yaml", ".yml"]
self.file_types = ["json", "yaml", "yml"]
# Replace default value with actual value
if "default" in value:
self.value = value["default"]
if key == "headers":
self.value = """{'Authorization':
'Bearer <token>'}"""
# Add options to openai
if name == "OpenAI" and key == "model_name":
self.options = constants.OPENAI_MODELS
self.is_list = True
elif name == "ChatOpenAI" and key == "model_name":
self.options = constants.CHAT_OPENAI_MODELS
self.is_list = True
class TemplateField(TemplateFieldCreator): class TemplateField(TemplateFieldCreator):
pass pass
@ -139,10 +70,10 @@ class Template(BaseModel):
class FrontendNode(BaseModel): class FrontendNode(BaseModel):
template: Template template: Template
description: str description: str
base_classes: list base_classes: List[str]
name: str = "" name: str = ""
def to_dict(self): def to_dict(self) -> dict:
return { return {
self.name: { self.name: {
"template": self.template.to_dict(self.format_field), "template": self.template.to_dict(self.format_field),
@ -153,53 +84,145 @@ class FrontendNode(BaseModel):
@staticmethod @staticmethod
def format_field(field: TemplateField, name: Optional[str] = None) -> None: def format_field(field: TemplateField, name: Optional[str] = None) -> None:
"""Formats a given field based on its attributes and value."""
SPECIAL_FIELD_HANDLERS = {
"allowed_tools": lambda field: "Tool",
"max_value_length": lambda field: "int",
}
key = field.name key = field.name
value = field.to_dict() value = field.to_dict()
_type = value["type"] _type = value["type"]
# Remove 'Optional' wrapper _type = FrontendNode.remove_optional(_type)
if "Optional" in _type: _type, is_list = FrontendNode.check_for_list_type(_type)
_type = _type.replace("Optional[", "")[:-1] field.is_list = is_list or field.is_list
_type = FrontendNode.replace_mapping_with_dict(_type)
_type = FrontendNode.handle_union_type(_type)
# Check for list type field.field_type = FrontendNode.handle_special_field(
if "List" in _type or "Sequence" in _type: field, key, _type, SPECIAL_FIELD_HANDLERS
_type = _type.replace("List[", "") )
_type = _type.replace("Sequence[", "")[:-1] field.field_type = FrontendNode.handle_dict_type(field, _type)
field.is_list = True field.show = FrontendNode.should_show_field(key, field.required)
field.password = FrontendNode.should_be_password(key, field.show)
field.multiline = FrontendNode.should_be_multiline(key)
# Replace 'Mapping' with 'dict' FrontendNode.replace_default_value(field, value)
if "Mapping" in _type: FrontendNode.handle_specific_field_values(field, key, name)
_type = _type.replace("Mapping", "dict") FrontendNode.handle_kwargs_field(field)
FrontendNode.handle_api_key_field(field, key)
# {'type': 'Union[float, Tuple[float, float], NoneType]'} != {'type': 'float'} @staticmethod
def remove_optional(_type: str) -> str:
"""Removes 'Optional' wrapper from the type if present."""
return re.sub(r"Optional\[(.*)\]", r"\1", _type)
@staticmethod
def check_for_list_type(_type: str) -> tuple:
"""Checks for list type and returns the modified type and a boolean indicating if it's a list."""
is_list = "List" in _type or "Sequence" in _type
if is_list:
_type = re.sub(r"(List|Sequence)\[(.*)\]", r"\2", _type)
return _type, is_list
@staticmethod
def replace_mapping_with_dict(_type: str) -> str:
"""Replaces 'Mapping' with 'dict'."""
return _type.replace("Mapping", "dict")
@staticmethod
def handle_union_type(_type: str) -> str:
"""Simplifies the 'Union' type to the first type in the Union."""
if "Union" in _type: if "Union" in _type:
_type = _type.replace("Union[", "")[:-1] _type = _type.replace("Union[", "")[:-1]
_type = _type.split(",")[0] _type = _type.split(",")[0]
_type = _type.replace("]", "").replace("[", "") _type = _type.replace("]", "").replace("[", "")
return _type
field.field_type = _type @staticmethod
def handle_special_field(
field, key: str, _type: str, SPECIAL_FIELD_HANDLERS
) -> str:
"""Handles special field by using the respective handler if present."""
handler = SPECIAL_FIELD_HANDLERS.get(key)
return handler(field) if handler else _type
# Change type from str to Tool @staticmethod
field.field_type = "Tool" if key in {"allowed_tools"} else field.field_type def handle_dict_type(field: TemplateField, _type: str) -> str:
"""Handles 'dict' type by replacing it with 'code' or 'file' based on the field name."""
if "dict" in _type.lower():
if field.name == "dict_":
field.field_type = "file"
field.suffixes = [".json", ".yaml", ".yml"]
field.file_types = ["json", "yaml", "yml"]
else:
field.field_type = "code"
return _type
field.field_type = "int" if key in {"max_value_length"} else field.field_type @staticmethod
def replace_default_value(field: TemplateField, value: dict) -> None:
"""Replaces default value with actual value if 'default' is present in value."""
if "default" in value:
field.value = value["default"]
# Show or not field @staticmethod
field.show = bool( def handle_specific_field_values(
(field.required and key not in ["input_variables"]) field: TemplateField, key: str, name: Optional[str] = None
) -> None:
"""Handles specific field values for certain fields."""
if key == "headers":
field.value = """{'Authorization':
'Bearer <token>'}"""
if name == "OpenAI" and key == "model_name":
field.options = constants.OPENAI_MODELS
field.is_list = True
elif name == "ChatOpenAI" and key == "model_name":
field.options = constants.CHAT_OPENAI_MODELS
field.is_list = True
if "api_key" in key and "OpenAI" in str(name):
field.display_name = "OpenAI API Key"
field.required = False
if field.value is None:
field.value = ""
@staticmethod
def handle_kwargs_field(field: TemplateField) -> None:
"""Handles kwargs field by setting certain attributes."""
if "kwargs" in field.name.lower():
field.advanced = True
field.required = False
field.show = False
@staticmethod
def handle_api_key_field(field: TemplateField, key: str) -> None:
"""Handles api key field by setting certain attributes."""
if "api" in key.lower() and "key" in key.lower():
field.required = False
field.advanced = False
@staticmethod
def should_show_field(key: str, required: bool) -> bool:
"""Determines whether the field should be shown."""
return (
(required and key not in ["input_variables"])
or key in FORCE_SHOW_FIELDS or key in FORCE_SHOW_FIELDS
or "api" in key or "api" in key
or ("key" in key and "input" not in key and "output" not in key) or ("key" in key and "input" not in key and "output" not in key)
) )
# Add password field @staticmethod
field.password = ( def should_be_password(key: str, show: bool) -> bool:
"""Determines whether the field should be a password field."""
return (
any(text in key.lower() for text in {"password", "token", "api", "key"}) any(text in key.lower() for text in {"password", "token", "api", "key"})
and field.show and show
) )
# Add multline @staticmethod
field.multiline = key in { def should_be_multiline(key: str) -> bool:
"""Determines whether the field should be multiline."""
return key in {
"suffix", "suffix",
"prefix", "prefix",
"template", "template",
@ -209,43 +232,24 @@ class FrontendNode(BaseModel):
"description", "description",
} }
# Replace dict type with str @staticmethod
if "dict" in field.field_type.lower(): def replace_dict_with_code_or_file(
field.field_type = "code" field: TemplateField, _type: str, key: str
) -> str:
"""Replaces 'dict' type with 'code' or 'file'."""
if "dict" in _type.lower():
if key == "dict_":
field.field_type = "file"
field.suffixes = [".json", ".yaml", ".yml"]
field.file_types = ["json", "yaml", "yml"]
else:
field.field_type = "code"
return field.field_type
if key == "dict_": @staticmethod
field.field_type = "file" def set_field_default_value(field: TemplateField, value: dict, key: str) -> None:
field.suffixes = [".json", ".yaml", ".yml"] """Sets the field value with the default value if present."""
field.file_types = ["json", "yaml", "yml"]
# Replace default value with actual value
if "default" in value: if "default" in value:
field.value = value["default"] field.value = value["default"]
if key == "headers": if key == "headers":
field.value = """{'Authorization': field.value = """{'Authorization': 'Bearer <token>'}"""
'Bearer <token>'}"""
# Add options to openai
if name == "OpenAI" and key == "model_name":
field.options = constants.OPENAI_MODELS
field.is_list = True
elif name == "ChatOpenAI":
if key == "model_name":
field.options = constants.CHAT_OPENAI_MODELS
field.is_list = True
if "api_key" in key and "OpenAI" in str(name):
field.display_name = "OpenAI API Key"
field.required = False
if field.value is None:
field.value = ""
if "kwargs" in field.name.lower():
field.advanced = True
field.required = False
field.show = False
# If the field.name contains api or api and key, then it might be an api key
# other conditions are to make sure that it is not an input or output variable
if "api" in key.lower() and "key" in key.lower():
field.required = False
field.advanced = False