✨ 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.directory_reader import DirectoryReader
|
||||||
from langflow.interface.custom.utils import (
|
from langflow.interface.custom.utils import (
|
||||||
extract_inner_type_from_generic_alias,
|
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.models.flow import Flow
|
||||||
from langflow.services.database.utils import session_getter
|
from langflow.services.database.utils import session_getter
|
||||||
from langflow.services.deps import get_credential_service, get_db_service
|
from langflow.services.deps import get_credential_service, get_db_service
|
||||||
|
|
@ -29,6 +28,7 @@ class CustomComponent(Component):
|
||||||
repr_value: Optional[Any] = ""
|
repr_value: Optional[Any] = ""
|
||||||
user_id: Optional[Union[UUID, str]] = None
|
user_id: Optional[Union[UUID, str]] = None
|
||||||
status: Optional[Any] = None
|
status: Optional[Any] = None
|
||||||
|
_tree: Optional[dict] = None
|
||||||
|
|
||||||
def __init__(self, **data):
|
def __init__(self, **data):
|
||||||
self.cache = TTLCache(maxsize=1024, ttl=60)
|
self.cache = TTLCache(maxsize=1024, ttl=60)
|
||||||
|
|
@ -78,8 +78,11 @@ class CustomComponent(Component):
|
||||||
def validate(self) -> bool:
|
def validate(self) -> bool:
|
||||||
return self._class_template_validation(self.code) if self.code else False
|
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
|
@property
|
||||||
def get_function_entrypoint_args(self) -> list:
|
def get_function_entrypoint_args(self) -> list:
|
||||||
|
|
@ -108,9 +111,10 @@ class CustomComponent(Component):
|
||||||
def get_build_method(self):
|
def get_build_method(self):
|
||||||
if not self.code:
|
if not self.code:
|
||||||
return []
|
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:
|
if not component_classes:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
@ -123,16 +127,19 @@ class CustomComponent(Component):
|
||||||
if not build_methods:
|
if not build_methods:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
return build_methods[0]
|
return build_methods[0]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def get_function_entrypoint_return_type(self) -> List[Any]:
|
def get_function_entrypoint_return_type(self) -> List[Any]:
|
||||||
build_method = self.get_build_method()
|
build_method = self.get_build_method()
|
||||||
if not build_method:
|
if not build_method:
|
||||||
return build_method
|
|
||||||
return_type = build_method["return_type"]
|
|
||||||
if not return_type:
|
|
||||||
return []
|
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 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]:
|
if hasattr(return_type, "__origin__") and return_type.__origin__ in [list, List]:
|
||||||
return_type = extract_inner_type_from_generic_alias(return_type)
|
return_type = extract_inner_type_from_generic_alias(return_type)
|
||||||
|
|
@ -152,13 +159,13 @@ class CustomComponent(Component):
|
||||||
def get_main_class_name(self):
|
def get_main_class_name(self):
|
||||||
if not self.code:
|
if not self.code:
|
||||||
return ""
|
return ""
|
||||||
tree = self.get_code_tree(self.code)
|
|
||||||
|
|
||||||
base_name = self.code_class_base_inheritance
|
base_name = self.code_class_base_inheritance
|
||||||
method_name = self.function_entrypoint_name
|
method_name = self.function_entrypoint_name
|
||||||
|
|
||||||
classes = []
|
classes = []
|
||||||
for item in tree.get("classes", []):
|
for item in self.tree.get("classes", []):
|
||||||
if base_name in item["bases"]:
|
if base_name in item["bases"]:
|
||||||
method_names = [method["name"] for method in item["methods"]]
|
method_names = [method["name"] for method in item["methods"]]
|
||||||
if method_name in method_names:
|
if method_name in method_names:
|
||||||
|
|
@ -171,11 +178,11 @@ class CustomComponent(Component):
|
||||||
def build_template_config(self):
|
def build_template_config(self):
|
||||||
if not self.code:
|
if not self.code:
|
||||||
return {}
|
return {}
|
||||||
tree = self.get_code_tree(self.code)
|
|
||||||
|
|
||||||
attributes = [
|
attributes = [
|
||||||
main_class["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
|
if main_class["name"] == self.get_main_class_name
|
||||||
]
|
]
|
||||||
# Get just the first item
|
# Get just the first item
|
||||||
|
|
@ -219,7 +226,8 @@ class CustomComponent(Component):
|
||||||
return validate.create_function(self.code, self.function_entrypoint_name)
|
return validate.create_function(self.code, self.function_entrypoint_name)
|
||||||
|
|
||||||
async def load_flow(self, flow_id: str, tweaks: Optional[dict] = None) -> Any:
|
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()
|
db_service = get_db_service()
|
||||||
with session_getter(db_service) as session:
|
with session_getter(db_service) as session:
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue