Refactor tools and constants modules

- Refactored the tools module to improve code readability and maintainability.
- Fixed import statements in the constants module to correctly import modules.
- Updated documentation and added comments for better understanding of the code.
This commit is contained in:
gustavoschaedler 2023-07-05 19:57:28 +01:00
commit ecbb0feeab
5 changed files with 28 additions and 45 deletions

View file

@ -12,7 +12,7 @@ from langflow.interface.importing.utils import import_class
from langflow.interface.tools.custom import ( from langflow.interface.tools.custom import (
PythonFunctionTool, PythonFunctionTool,
PythonFunction, PythonFunction,
CustomComponent CustomComponent,
) )
FILE_TOOLS = {"JsonSpec": JsonSpec} FILE_TOOLS = {"JsonSpec": JsonSpec}
@ -23,8 +23,7 @@ CUSTOM_TOOLS = {
"PythonFunction": PythonFunction, "PythonFunction": PythonFunction,
} }
OTHER_TOOLS = {tool: import_class(f"langchain.tools.{tool}") OTHER_TOOLS = {tool: import_class(f"langchain.tools.{tool}") for tool in tools.__all__}
for tool in tools.__all__}
ALL_TOOLS_NAMES = { ALL_TOOLS_NAMES = {
**_BASE_TOOLS, **_BASE_TOOLS,

View file

@ -85,19 +85,11 @@ class CustomComponent(BaseModel):
code: str code: str
function: Optional[Callable] = None function: Optional[Callable] = None
function_entrypoint_name = "build" function_entrypoint_name = "build"
return_type_valid_list = [ return_type_valid_list = ["ConversationChain", "BaseLLM", "Tool"]
"ConversationChain",
"BaseLLM",
"Tool"
]
class_template = { class_template = {
"imports": [], "imports": [],
"class": { "class": {"inherited_classes": "", "name": "", "init": ""},
"inherited_classes": "", "functions": [],
"name": "",
"init": ""
},
"functions": []
} }
def __init__(self, **data): def __init__(self, **data):
@ -105,15 +97,18 @@ class CustomComponent(BaseModel):
def _handle_import(self, node): def _handle_import(self, node):
for alias in node.names: for alias in node.names:
module_name = getattr(node, 'module', None) module_name = getattr(node, "module", None)
self.class_template['imports'].append( self.class_template["imports"].append(
f"{module_name}.{alias.name}" if module_name else alias.name) f"{module_name}.{alias.name}" if module_name else alias.name
)
def _handle_class(self, node): def _handle_class(self, node):
self.class_template['class'].update({ self.class_template["class"].update(
'name': node.name, {
'inherited_classes': [ast.unparse(base) for base in node.bases] "name": node.name,
}) "inherited_classes": [ast.unparse(base) for base in node.bases],
}
)
for inner_node in node.body: for inner_node in node.body:
if isinstance(inner_node, ast.FunctionDef): if isinstance(inner_node, ast.FunctionDef):
@ -122,20 +117,20 @@ class CustomComponent(BaseModel):
def _handle_function(self, node): def _handle_function(self, node):
function_name = node.name function_name = node.name
function_args_str = ast.unparse(node.args) function_args_str = ast.unparse(node.args)
function_args = function_args_str.split( function_args = function_args_str.split(", ") if function_args_str else []
", ") if function_args_str else []
return_type = ast.unparse(node.returns) if node.returns else "None" return_type = ast.unparse(node.returns) if node.returns else "None"
function_data = { function_data = {
"name": function_name, "name": function_name,
"arguments": function_args, "arguments": function_args,
"return_type": return_type "return_type": return_type,
} }
if function_name == "__init__": if function_name == "__init__":
self.class_template['class']['init'] = function_args_str.split( self.class_template["class"]["init"] = (
", ") if function_args_str else [] function_args_str.split(", ") if function_args_str else []
)
else: else:
self.class_template["functions"].append(function_data) self.class_template["functions"].append(function_data)
@ -143,7 +138,7 @@ class CustomComponent(BaseModel):
output_list = [] output_list = []
for item in input_list: for item in input_list:
# Split each item on ':' to separate variable name and type # Split each item on ':' to separate variable name and type
split_item = item.split(':') split_item = item.split(":")
# If there is a type, strip any leading/trailing spaces from it # If there is a type, strip any leading/trailing spaces from it
if len(split_item) > 1: if len(split_item) > 1:
@ -161,10 +156,7 @@ class CustomComponent(BaseModel):
except SyntaxError as err: except SyntaxError as err:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
detail={ detail={"error": err.msg, "traceback": traceback.format_exc()},
'error': err.msg,
'traceback': traceback.format_exc()
},
) from err ) from err
for node in module.body: for node in module.body:
@ -180,8 +172,7 @@ class CustomComponent(BaseModel):
functions = data.get("functions", []) functions = data.get("functions", [])
if build_function := next( if build_function := next(
(f for f in functions if f["name"] (f for f in functions if f["name"] == self.function_entrypoint_name),
== self.function_entrypoint_name),
None, None,
): ):
function_args = build_function.get("arguments", None) function_args = build_function.get("arguments", None)
@ -201,8 +192,7 @@ class CustomComponent(BaseModel):
functions = code.get("functions", []) functions = code.get("functions", [])
if build_function := next( if build_function := next(
(f for f in functions if f["name"] (f for f in functions if f["name"] == self.function_entrypoint_name),
== self.function_entrypoint_name),
None, None,
): ):
# Check if the return type of the build function is valid # Check if the return type of the build function is valid
@ -211,10 +201,7 @@ class CustomComponent(BaseModel):
return False return False
def get_function(self): def get_function(self):
return validate.create_function( return validate.create_function(self.code, self.function_entrypoint_name)
self.code,
self.function_entrypoint_name
)
@property @property
def data(self): def data(self):

View file

@ -3,7 +3,7 @@ from langflow.template.frontend_node.base import FrontendNode
from langflow.template.template.base import Template from langflow.template.template.base import Template
from langflow.utils.constants import ( from langflow.utils.constants import (
DEFAULT_PYTHON_FUNCTION, DEFAULT_PYTHON_FUNCTION,
DEFAULT_CUSTOM_COMPONENT_CODE DEFAULT_CUSTOM_COMPONENT_CODE,
) )

View file

@ -25,7 +25,6 @@ ANTHROPIC_MODELS = [
"claude-instant-v1", "claude-instant-v1",
# Like claude-instant-v1 with a 100,000 token context window but retains its performance. # Like claude-instant-v1 with a 100,000 token context window but retains its performance.
"claude-instant-v1-100k", "claude-instant-v1-100k",
# Specific sub-versions of the above models: # Specific sub-versions of the above models:
# Vs claude-v1.2: better instruction-following, code, and non-English dialogue and writing. # Vs claude-v1.2: better instruction-following, code, and non-English dialogue and writing.
"claude-v1.3", "claude-v1.3",

View file

@ -108,8 +108,7 @@ def execute_function(code, function_name, *args, **kwargs):
try: try:
exec(code_obj, exec_globals, locals()) exec(code_obj, exec_globals, locals())
except Exception as exc: except Exception as exc:
raise ValueError( raise ValueError("Function string does not contain a function") from exc
"Function string does not contain a function") from exc
# Add the function to the exec_globals dictionary # Add the function to the exec_globals dictionary
exec_globals[function_name] = locals()[function_name] exec_globals[function_name] = locals()[function_name]
@ -190,8 +189,7 @@ def create_class(code, class_name):
try: try:
imported_module = importlib.import_module(node.module) imported_module = importlib.import_module(node.module)
for alias in node.names: for alias in node.names:
exec_globals[alias.name] = getattr( exec_globals[alias.name] = getattr(imported_module, alias.name)
imported_module, alias.name)
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
raise ModuleNotFoundError( raise ModuleNotFoundError(
f"Module {node.module} not found. Please install it and try again." f"Module {node.module} not found. Please install it and try again."