From b1224ddb63c9a81a0340c4eb74438b8ecd10a9e5 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Fri, 1 Dec 2023 13:12:48 -0300 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20feat(custom=5Fcomponent.py):=20add?= =?UTF-8?q?=20=5Ftree=20attribute=20to=20store=20code=20tree=20for=20bette?= =?UTF-8?q?r=20performance=20=F0=9F=94=A7=20refactor(custom=5Fcomponent.py?= =?UTF-8?q?):=20refactor=20get=5Fcode=5Ftree=20method=20to=20use=20the=20?= =?UTF-8?q?=5Ftree=20attribute=20=F0=9F=94=A7=20refactor(custom=5Fcomponen?= =?UTF-8?q?t.py):=20refactor=20get=5Fbuild=5Fmethod=20method=20to=20use=20?= =?UTF-8?q?the=20tree=20attribute=20=F0=9F=94=A7=20refactor(custom=5Fcompo?= =?UTF-8?q?nent.py):=20refactor=20get=5Fmain=5Fclass=5Fname=20method=20to?= =?UTF-8?q?=20use=20the=20tree=20attribute=20=F0=9F=94=A7=20refactor(custo?= =?UTF-8?q?m=5Fcomponent.py):=20refactor=20build=5Ftemplate=5Fconfig=20met?= =?UTF-8?q?hod=20to=20use=20the=20tree=20attribute=20=F0=9F=94=A7=20refact?= =?UTF-8?q?or(custom=5Fcomponent.py):=20refactor=20load=5Fflow=20method=20?= =?UTF-8?q?to=20use=20import=20statement=20on=20separate=20lines=20for=20b?= =?UTF-8?q?etter=20readability?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../interface/custom/custom_component.py | 36 +++++++++++-------- 1 file changed, 22 insertions(+), 14 deletions(-) diff --git a/src/backend/langflow/interface/custom/custom_component.py b/src/backend/langflow/interface/custom/custom_component.py index 1ea2a1c78..14a0246b7 100644 --- a/src/backend/langflow/interface/custom/custom_component.py +++ b/src/backend/langflow/interface/custom/custom_component.py @@ -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: