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: