✨ feat(custom_component.py): add _tree attribute to store code tree for better performance
🔧 refactor(custom_component.py): refactor get_code_tree method to use the _tree attribute 🔧 refactor(custom_component.py): refactor get_build_method method to use the tree attribute 🔧 refactor(custom_component.py): refactor get_main_class_name method to use the tree attribute 🔧 refactor(custom_component.py): refactor build_template_config method to use the tree attribute 🔧 refactor(custom_component.py): refactor load_flow method to use import statement on separate lines for better readability
This commit is contained in:
parent
7a2ba97b4c
commit
b1224ddb63
1 changed files with 22 additions and 14 deletions
|
|
@ -10,8 +10,7 @@ from langflow.interface.custom.component import Component
|
|||
from langflow.interface.custom.directory_reader import DirectoryReader
|
||||
from langflow.interface.custom.utils import (
|
||||
extract_inner_type_from_generic_alias,
|
||||
extract_union_types_from_generic_alias,
|
||||
)
|
||||
extract_union_types_from_generic_alias)
|
||||
from langflow.services.database.models.flow import Flow
|
||||
from langflow.services.database.utils import session_getter
|
||||
from langflow.services.deps import get_credential_service, get_db_service
|
||||
|
|
@ -29,6 +28,7 @@ class CustomComponent(Component):
|
|||
repr_value: Optional[Any] = ""
|
||||
user_id: Optional[Union[UUID, str]] = None
|
||||
status: Optional[Any] = None
|
||||
_tree: Optional[dict] = None
|
||||
|
||||
def __init__(self, **data):
|
||||
self.cache = TTLCache(maxsize=1024, ttl=60)
|
||||
|
|
@ -78,8 +78,11 @@ class CustomComponent(Component):
|
|||
def validate(self) -> bool:
|
||||
return self._class_template_validation(self.code) if self.code else False
|
||||
|
||||
def get_code_tree(self, code: str):
|
||||
return super().get_code_tree(code)
|
||||
|
||||
|
||||
@property
|
||||
def tree(self):
|
||||
return self.get_code_tree(self.code)
|
||||
|
||||
@property
|
||||
def get_function_entrypoint_args(self) -> list:
|
||||
|
|
@ -108,9 +111,10 @@ class CustomComponent(Component):
|
|||
def get_build_method(self):
|
||||
if not self.code:
|
||||
return []
|
||||
tree = self.get_code_tree(self.code)
|
||||
|
||||
component_classes = [cls for cls in tree["classes"] if self.code_class_base_inheritance in cls["bases"]]
|
||||
|
||||
|
||||
component_classes = [cls for cls in self.tree["classes"] if self.code_class_base_inheritance in cls["bases"]]
|
||||
if not component_classes:
|
||||
return []
|
||||
|
||||
|
|
@ -123,16 +127,19 @@ class CustomComponent(Component):
|
|||
if not build_methods:
|
||||
return []
|
||||
|
||||
|
||||
return build_methods[0]
|
||||
|
||||
@property
|
||||
def get_function_entrypoint_return_type(self) -> List[Any]:
|
||||
build_method = self.get_build_method()
|
||||
if not build_method:
|
||||
return build_method
|
||||
return_type = build_method["return_type"]
|
||||
if not return_type:
|
||||
return []
|
||||
elif not build_method["has_return"]:
|
||||
return []
|
||||
|
||||
return_type = build_method["return_type"]
|
||||
|
||||
# If list or List is in the return type, then we remove it and return the inner type
|
||||
if hasattr(return_type, "__origin__") and return_type.__origin__ in [list, List]:
|
||||
return_type = extract_inner_type_from_generic_alias(return_type)
|
||||
|
|
@ -152,13 +159,13 @@ class CustomComponent(Component):
|
|||
def get_main_class_name(self):
|
||||
if not self.code:
|
||||
return ""
|
||||
tree = self.get_code_tree(self.code)
|
||||
|
||||
|
||||
base_name = self.code_class_base_inheritance
|
||||
method_name = self.function_entrypoint_name
|
||||
|
||||
classes = []
|
||||
for item in tree.get("classes", []):
|
||||
for item in self.tree.get("classes", []):
|
||||
if base_name in item["bases"]:
|
||||
method_names = [method["name"] for method in item["methods"]]
|
||||
if method_name in method_names:
|
||||
|
|
@ -171,11 +178,11 @@ class CustomComponent(Component):
|
|||
def build_template_config(self):
|
||||
if not self.code:
|
||||
return {}
|
||||
tree = self.get_code_tree(self.code)
|
||||
|
||||
|
||||
attributes = [
|
||||
main_class["attributes"]
|
||||
for main_class in tree.get("classes", [])
|
||||
for main_class in self.tree.get("classes", [])
|
||||
if main_class["name"] == self.get_main_class_name
|
||||
]
|
||||
# Get just the first item
|
||||
|
|
@ -219,7 +226,8 @@ class CustomComponent(Component):
|
|||
return validate.create_function(self.code, self.function_entrypoint_name)
|
||||
|
||||
async def load_flow(self, flow_id: str, tweaks: Optional[dict] = None) -> Any:
|
||||
from langflow.processing.process import build_sorted_vertices, process_tweaks
|
||||
from langflow.processing.process import (build_sorted_vertices,
|
||||
process_tweaks)
|
||||
|
||||
db_service = get_db_service()
|
||||
with session_getter(db_service) as session:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue