fix: removing dicts from inside class to stop recreating it

This commit is contained in:
Gabriel Almeida 2023-03-31 14:01:35 -03:00
commit 0858734eb0

View file

@ -2,9 +2,10 @@ from langflow.custom import customs
from langflow.interface.tools.constants import ( from langflow.interface.tools.constants import (
ALL_TOOLS_NAMES, ALL_TOOLS_NAMES,
CUSTOM_TOOLS, CUSTOM_TOOLS,
OTHER_TOOLS, FILE_TOOLS,
) )
from langflow.template.template import Field, Template from langflow.template.base import Field
from langflow.template.base import Template
from langflow.utils import util from langflow.utils import util
from langflow.settings import settings from langflow.settings import settings
from langflow.interface.base import LangChainTypeCreator from langflow.interface.base import LangChainTypeCreator
@ -22,32 +23,7 @@ from langflow.interface.tools.util import (
) )
class ToolCreator(LangChainTypeCreator): TOOL_INPUTS = {
type_name: str = "tools"
tools_dict: Dict | None = None
@property
def type_to_loader_dict(self) -> Dict:
if self.tools_dict is None:
self.tools_dict = get_tools_dict()
return self.tools_dict
def get_signature(self, name: str) -> Dict | None:
"""Get the signature of a tool."""
NODE_INPUTS = ["llm", "func"]
base_classes = ["Tool"]
all_tools = {}
for tool in self.type_to_loader_dict.keys():
if tool_params := get_tool_params(get_tool_by_name(tool)):
tool_name = tool_params.get("name") or str(tool)
all_tools[tool_name] = {"type": tool, "params": tool_params}
# Raise error if name is not in tools
if name not in all_tools.keys():
raise ValueError("Tool not found")
type_dict = {
"str": Field( "str": Field(
field_type="str", field_type="str",
required=True, required=True,
@ -79,7 +55,32 @@ class ToolCreator(LangChainTypeCreator):
show=True, show=True,
value="", value="",
), ),
} }
class ToolCreator(LangChainTypeCreator):
type_name: str = "tools"
tools_dict: Dict | None = None
@property
def type_to_loader_dict(self) -> Dict:
if self.tools_dict is None:
self.tools_dict = get_tools_dict()
return self.tools_dict
def get_signature(self, name: str) -> Dict | None:
"""Get the signature of a tool."""
base_classes = ["Tool"]
all_tools = {}
for tool in self.type_to_loader_dict.keys():
if tool_params := get_tool_params(get_tool_by_name(tool)):
tool_name = tool_params.get("name") or str(tool)
all_tools[tool_name] = {"type": tool, "params": tool_params}
# Raise error if name is not in tools
if name not in all_tools.keys():
raise ValueError("Tool not found")
tool_type: str = all_tools[name]["type"] # type: ignore tool_type: str = all_tools[name]["type"] # type: ignore
@ -101,8 +102,9 @@ class ToolCreator(LangChainTypeCreator):
base_classes = ["function"] base_classes = ["function"]
if node := customs.get_custom_nodes("tools").get(tool_type): if node := customs.get_custom_nodes("tools").get(tool_type):
return node return node
elif tool_type in OTHER_TOOLS: elif tool_type in FILE_TOOLS:
params = all_tools[name]["params"] # type: ignore params = all_tools[name]["params"] # type: ignore
base_classes += [name]
else: else:
params = [] params = []
@ -110,10 +112,7 @@ class ToolCreator(LangChainTypeCreator):
# Copy the field and add the name # Copy the field and add the name
fields = [] fields = []
for param in params: for param in params:
if param in NODE_INPUTS: field = TOOL_INPUTS.get(param, TOOL_INPUTS["str"])
field = type_dict[param].copy()
else:
field = type_dict["str"].copy()
field.name = param field.name = param
if param == "aiosession": if param == "aiosession":
field.show = False field.show = False
@ -122,9 +121,7 @@ class ToolCreator(LangChainTypeCreator):
template = Template(fields=fields, type_name=tool_type) template = Template(fields=fields, type_name=tool_type)
tool_params = get_tool_params(get_tool_by_name(tool_type)) tool_params = all_tools[name]["params"]
if tool_params is None:
tool_params = {}
return { return {
"template": util.format_dict(template.to_dict()), "template": util.format_dict(template.to_dict()),
**tool_params, **tool_params,