feat: implementation of PythonFunction and class based templates

This commit is contained in:
Gabriel Almeida 2023-03-26 21:05:11 -03:00
commit a788d93682
13 changed files with 432 additions and 222 deletions

View file

@ -20,6 +20,8 @@ tools:
- PAL-MATH
- Calculator
- Serper Search
- Tool
- PythonFunction
memories:
# - ConversationBufferMemory

View file

@ -1,6 +1,12 @@
from langflow.node.nodes import ZeroShotPromptNode
from langflow.node import nodes
def get_custom_prompts():
"""Get custom prompts."""
return ZeroShotPromptNode().to_dict()
CUSTOM_NODES = {
"prompts": {**nodes.ZeroShotPromptNode().to_dict()},
"tools": {**nodes.PythonFunctionNode().to_dict(), **nodes.ToolNode().to_dict()},
}
def get_custom_nodes(node_type: str):
"""Get custom nodes."""
return CUSTOM_NODES.get(node_type, [])

View file

@ -0,0 +1,13 @@
from pydantic import BaseModel, validator
from langchain.agents import tool
class PythonFunction(BaseModel):
code: str
# Validate the function
@validator("code")
def validate_func(cls, v):
# Validate with LangChain's tool decorator
tool(v)
return v

View file

@ -8,7 +8,7 @@ from langchain.agents import Agent
from langchain.chains.base import Chain
from langchain.llms.base import BaseLLM
from langchain.tools import BaseTool
from langflow.utils.util import get_tools_dict
from langflow.utils.util import get_tool_by_name
def import_module(module_path: str) -> Any:
@ -55,7 +55,7 @@ def import_llm(llm: str) -> BaseLLM:
def import_tool(tool: str) -> BaseTool:
"""Import tool from tool name"""
return get_tools_dict(tool)
return get_tool_by_name(tool)
def import_chain(chain: str) -> Chain:

View file

@ -17,13 +17,13 @@ def get_type_dict():
"prompts": list_prompts,
"llms": list_llms,
"tools": list_tools,
# "memories": list_memories,
"memories": list_memories,
}
def list_type(object_type: str):
"""List all components"""
return get_type_dict().get(object_type, lambda: "Invalid type")()
return get_type_dict().get(object_type, lambda: None)()
def list_agents():
@ -37,7 +37,7 @@ def list_agents():
def list_prompts():
"""List all prompt types"""
custom_prompts = customs.get_custom_prompts()
custom_prompts = customs.get_custom_nodes("prompts")
library_prompts = [
prompt.__annotations__["return"].__name__
for prompt in prompts.loading.type_to_loader_dict.values()
@ -52,11 +52,13 @@ def list_tools():
tools = []
for tool in get_all_tool_names():
tool_params = util.get_tool_params(util.get_tools_dict(tool))
tool_params = util.get_tool_params(util.get_tool_by_name(tool))
if tool_params and tool_params["name"] in settings.tools or settings.dev:
tools.append(tool_params["name"])
return tools
# Add Tool
custom_tools = customs.get_custom_nodes("tools")
return tools + list(custom_tools.keys())
def list_llms():

View file

@ -41,7 +41,7 @@ def instantiate_class(node_type: str, base_type: str, params: Dict) -> Any:
# which will be a str containing a python function
# and then we need to compile it and return the function
# as the instance
function_string = params["function"]
function_string = params["code"]
if isinstance(function_string, str):
return util.eval_function(function_string)
raise ValueError("Function should be a string")

View file

@ -15,6 +15,7 @@ from langflow.interface.custom_lists import (
memory_type_to_cls_dict,
)
from langflow.utils import util
from langflow.utils.constants import CUSTOM_TOOLS
def get_signature(name: str, object_type: str):
@ -53,8 +54,8 @@ def get_agent_signature(name: str):
def get_prompt_signature(name: str):
"""Get the signature of a prompt."""
try:
if name in customs.get_custom_prompts().keys():
return customs.get_custom_prompts()[name]
if name in customs.get_custom_nodes("prompts").keys():
return customs.get_custom_nodes("prompts")[name]
return util.build_template_from_function(
name, prompts.loading.type_to_loader_dict
)
@ -82,12 +83,14 @@ def get_tool_signature(name: str):
"""Get the signature of a tool."""
NODE_INPUTS = ["llm", "func"]
base_classes = ["Tool"]
all_tools = {}
for tool in get_all_tool_names():
if tool_params := util.get_tool_params(util.get_tools_dict(tool)):
all_tools[tool_params["name"]] = tool
all_tool_names = get_all_tool_names() + list(CUSTOM_TOOLS.keys())
for tool in all_tool_names:
if tool_params := util.get_tool_params(util.get_tool_by_name(tool)):
tool_name = tool_params.get("name") or str(tool)
all_tools[tool_name] = {"type": tool, "params": tool_params}
all_tools["BaseTool"] = "BaseTool"
# Raise error if name is not in tools
if name not in all_tools.keys():
raise ValueError("Tool not found")
@ -110,9 +113,17 @@ def get_tool_signature(name: str):
"value": "",
"multiline": True,
},
"code": {
"type": "str",
"required": True,
"list": False,
"show": True,
"value": "",
"multiline": True,
},
}
tool_type = all_tools[name]
tool_type = all_tools[name]["type"]
if tool_type in _BASE_TOOLS:
params = []
@ -124,8 +135,13 @@ def get_tool_signature(name: str):
elif tool_type in _EXTRA_OPTIONAL_TOOLS:
_, extra_keys = _EXTRA_OPTIONAL_TOOLS[tool_type]
params = extra_keys
elif tool_type == "BaseTool":
elif tool_type == "Tool":
params = ["name", "description", "func"]
elif tool_type in CUSTOM_TOOLS:
# Get custom tool params
params = all_tools[name]["params"]
base_classes = ["function"]
else:
params = []
@ -142,9 +158,8 @@ def get_tool_signature(name: str):
template["aiosession"]["show"] = False
template["_type"] = tool_type # type: ignore
return {
"template": template,
**util.get_tool_params(util.get_tools_dict(tool_type)),
"base_classes": ["Tool"],
"template": util.format_dict(template),
**util.get_tool_params(util.get_tool_by_name(tool_type)),
"base_classes": base_classes,
}

View file

@ -1,9 +1,10 @@
from langflow.node.template import Field, FrontendNode, Template
from langchain.agents.mrkl import prompt
from langflow.utils.constants import DEFAULT_PYTHON_FUNCTION
class ZeroShotPromptNode(FrontendNode):
_name = "ZeroShotPrompt"
name = "ZeroShotPrompt"
template = Template(
type_name="zero_shot",
fields=[
@ -44,3 +45,71 @@ class ZeroShotPromptNode(FrontendNode):
def to_dict(self):
return super().to_dict()
class PythonFunctionNode(FrontendNode):
name = "PythonFunction"
template = Template(
type_name="python_function",
fields=[
Field(
field_type="str",
required=True,
placeholder="",
is_list=False,
show=True,
multiline=True,
value=DEFAULT_PYTHON_FUNCTION,
name="code",
),
],
)
description = "Python function to be executed."
base_classes = ["function"]
def to_dict(self):
return super().to_dict()
class ToolNode(FrontendNode):
name = "Tool"
template = Template(
type_name="tool",
fields=[
Field(
field_type="str",
required=True,
placeholder="",
is_list=False,
show=True,
multiline=True,
value="",
name="name",
),
Field(
field_type="str",
required=True,
placeholder="",
is_list=False,
show=True,
multiline=True,
value="",
name="description",
),
Field(
field_type="str",
required=True,
placeholder="",
is_list=False,
show=True,
multiline=True,
value="",
name="func",
),
],
)
description = "Tool to be used in the flow."
base_classes = ["BaseTool"]
def to_dict(self):
return super().to_dict()

View file

@ -15,7 +15,9 @@ class Field(BaseModel):
name: str = None
def to_dict(self):
return self.dict()
result = self.dict()
result["type"] = result.pop("field_type")
return result
class Template(BaseModel):
@ -32,11 +34,11 @@ class FrontendNode(BaseModel):
template: Template
description: str
base_classes: list
_name: str = None
name: str = None
def to_dict(self):
return {
self._name: {
self.name: {
"template": self.template.to_dict(),
"description": self.description,
"base_classes": self.base_classes,

View file

@ -1,4 +1,5 @@
from langchain.agents import Tool
from langflow.interface.custom_types import PythonFunction
OPENAI_MODELS = [
"text-davinci-003",
@ -9,3 +10,9 @@ OPENAI_MODELS = [
]
CHAT_OPENAI_MODELS = ["gpt-3.5-turbo", "gpt-4", "gpt-4-32k"]
CUSTOM_TOOLS = {"Tool": Tool, "PythonFunction": PythonFunction}
DEFAULT_PYTHON_FUNCTION = """
def python_function(text: str) -> str:
return text
"""

View file

@ -14,7 +14,6 @@ from langchain.agents.load_tools import (
from langchain.agents.tools import Tool
from langchain.tools import BaseTool
from langflow.utils import constants
@ -164,42 +163,38 @@ def get_default_factory(module: str, function: str):
return None
class GenericTool(Tool):
"""Base class for all tools."""
def default_func(self, **kwargs):
"""Default function for the tool."""
return "Default function"
def __init__(
self,
name: str = "Tool name",
description: str = "Tool description",
func: callable = None,
):
"""Initialize the tool."""
super().__init__(name=name, description=description, func=func)
def get_base_tool(name, description, func: callable) -> BaseTool:
return GenericTool(func=func, name="Generic Tool", description="Bacon")
def get_tools_dict(name: Optional[str] = None):
def get_tools_dict():
"""Get the tools dictionary."""
tools = {
**_BASE_TOOLS,
**_LLM_TOOLS, # type: ignore
**{k: v[0] for k, v in _EXTRA_LLM_TOOLS.items()}, # type: ignore
**_LLM_TOOLS,
**{k: v[0] for k, v in _EXTRA_LLM_TOOLS.items()},
**{k: v[0] for k, v in _EXTRA_OPTIONAL_TOOLS.items()},
**constants.CUSTOM_TOOLS,
}
tools.update({"BaseTool": get_base_tool})
return tools[name] if name else tools
return tools
def get_tool_params(func, **kwargs):
def get_tool_by_name(name: str):
"""Get a tool from the tools dictionary."""
tools = get_tools_dict()
if name not in tools:
raise ValueError(f"{name} not found.")
return tools[name]
def get_tool_params(tool, **kwargs):
# Parse the function code into an abstract syntax tree
# Define if it is a function or a class
if inspect.isfunction(tool):
return get_func_tool_params(tool, **kwargs)
elif inspect.isclass(tool):
# Get the parameters necessary to
# instantiate the class
return get_class_tool_params(tool, **kwargs)
def get_func_tool_params(func, **kwargs):
tree = ast.parse(inspect.getsource(func))
# Iterate over the statements in the abstract syntax tree
@ -208,7 +203,7 @@ def get_tool_params(func, **kwargs):
if isinstance(node, ast.Return):
tool = node.value
if isinstance(tool, ast.Call):
if tool.func.id == "Tool":
if isinstance(tool.func, ast.Name) and tool.func.id == "Tool":
if tool.keywords:
tool_params = {}
for keyword in tool.keywords:
@ -224,6 +219,7 @@ def get_tool_params(func, **kwargs):
"name": ast.literal_eval(tool.args[0]),
"description": ast.literal_eval(tool.args[2]),
}
#
else:
# get the class object from the return statement
try:
@ -237,11 +233,42 @@ def get_tool_params(func, **kwargs):
"name": getattr(class_obj, "name"),
"description": getattr(class_obj, "description"),
}
# Return None if no return statement was found
# Return None if no return statement was found
return None
def get_class_tool_params(cls, **kwargs):
tree = ast.parse(inspect.getsource(cls))
tool_params = {}
# Iterate over the statements in the abstract syntax tree
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef):
# Find the class definition and look for methods
for stmt in node.body:
if isinstance(stmt, ast.FunctionDef) and stmt.name == "__init__":
# There is no assignment statements in the __init__ method
# So we need to get the params from the function definition
for arg in stmt.args.args:
if arg.arg == "name":
# It should be the name of the class
tool_params[arg.arg] = cls.__name__
elif arg.arg == "self":
continue
# If there is not default value, set it to an empty string
else:
try:
tool_params[arg.arg] = ast.literal_eval(arg.annotation)
except ValueError:
tool_params[arg.arg] = ""
elif not cls == Tool and isinstance(stmt, ast.AnnAssign):
# Get the attribute name and the annotation
tool_params[stmt.target.id] = ""
return tool_params
def get_class_doc(class_name):
"""
Extracts information from the docstring of a given class.
@ -332,7 +359,7 @@ def format_dict(d, name: Optional[str] = None):
_type = _type.replace("Mapping", "dict")
# Change type from str to Tool
value["type"] = "Tool" if key in ["allowed_tools", "func"] else _type
value["type"] = "Tool" if key in ["allowed_tools"] else _type
# Show or not field
value["show"] = bool(
@ -351,11 +378,11 @@ def format_dict(d, name: Optional[str] = None):
# Add password field
value["password"] = any(
text in key for text in ["password", "token", "api", "key"]
text in key.lower() for text in ["password", "token", "api", "key"]
)
# Add multline
value["multiline"] = key in ["suffix", "prefix", "template", "examples"]
value["multiline"] = key in ["suffix", "prefix", "template", "examples", "code"]
# Replace default value with actual value
if "default" in value: