diff --git a/langflow/backend/customs.py b/langflow/backend/customs.py new file mode 100644 index 000000000..0cf129b52 --- /dev/null +++ b/langflow/backend/customs.py @@ -0,0 +1,40 @@ +from langchain.agents.mrkl import prompt + + +def get_custom_prompts(): + return { + "ZeroShotPrompt": { + "template": { + "_type": "zero_shot", + "prefix": { + "type": "str", + "required": False, + "placeholder": "", + "list": False, + "show": True, + "multiline": True, + "value": prompt.PREFIX, + }, + "suffix": { + "type": "str", + "required": True, + "placeholder": "", + "list": False, + "show": True, + "multiline": True, + "value": prompt.SUFFIX, + }, + "format_instructions": { + "type": "str", + "required": False, + "placeholder": "", + "list": False, + "show": True, + "multiline": True, + "value": prompt.FORMAT_INSTRUCTIONS, + }, + }, + "description": "Prompt template for Zero Shot Agent.", + "base_classes": ["BasePromptTemplate"], + } + } diff --git a/langflow/backend/endpoints.py b/langflow/backend/endpoints.py index 4acbade24..a1db8fedd 100644 --- a/langflow/backend/endpoints.py +++ b/langflow/backend/endpoints.py @@ -9,6 +9,7 @@ from langchain.llms.loading import load_llm_from_config from langchain.prompts.loading import load_prompt_from_config from typing import Any + # build router router = APIRouter() @@ -26,6 +27,10 @@ def get_type_list(): @router.get("/all") def get_all(): + # library_prompts = { + # prompt: signature.get_prompt(prompt) for prompt in list_endpoints.list_prompts() + # } + # custom_prompts = customs.get_custom_prompts() return { "chains": { chain: signature.get_chain(chain) for chain in list_endpoints.list_chains() @@ -33,6 +38,7 @@ def get_all(): "agents": { agent: signature.get_agent(agent) for agent in list_endpoints.list_agents() }, + # "prompts": {**library_prompts, **custom_prompts}, "prompts": { prompt: signature.get_prompt(prompt) for prompt in list_endpoints.list_prompts() @@ -67,8 +73,22 @@ def get_all(): @router.post("/predict") def get_load(data: dict[str, Any]): + # Get type list type_list = get_type_list() + # Substitute ZeroShotPromt with PromptTemplate + for node in data['nodes']: + if node["data"]["type"] == "ZeroShotPrompt": + # Build Prompt Template + tools = [ + tool + for tool in data['nodes'] + if tool["type"] != "chatOutputNode" + and "Tool" in tool["data"]["node"]["base_classes"] + ] + node["data"] = build_prompt_template(prompt=node["data"], tools=tools) + break + # Add input variables data = payload.extract_input_variables(data) @@ -96,12 +116,75 @@ def get_load(data: dict[str, Any]): else: return {"result": "Error: Type should be either agent, chain or llm"} - # elif extracted_json["_type"] in type_list["prompts"]: - # loaded = load_prompt_from_config(extracted_json) - # print(loaded.format(product='')) - # return {'result': loaded.format(product=message)} +def build_prompt_template(prompt, tools): + prefix = prompt["node"]["template"]["prefix"]["value"] + suffix = prompt["node"]["template"]["suffix"]["value"] + format_instructions = prompt["node"]["template"]["format_instructions"]["value"] - # if type in a["prompts"]: + tool_strings = "\n".join( + [ + f"{tool['data']['node']['name']}: {tool['data']['node']['description']}" + for tool in tools + ] + ) + tool_names = ", ".join([tool["data"]["node"]["name"] for tool in tools]) + format_instructions = format_instructions.format(tool_names=tool_names) + value = "\n\n".join([prefix, tool_strings, format_instructions, suffix]) - # return a + prompt["type"] = "PromptTemplate" + # prompt["value"] = value + + prompt["node"] = { + "template": { + "_type": "prompt", + "input_variables": { + "type": "str", + "required": True, + "placeholder": "", + "list": True, + "show": False, + "multiline": False, + }, + "output_parser": { + "type": "BaseOutputParser", + "required": False, + "placeholder": "", + "list": False, + "show": False, + "multline": False, + "value": None, + }, + "template": { + "type": "str", + "required": True, + "placeholder": "", + "list": False, + "show": True, + "multiline": True, + "value": value, + }, + "template_format": { + "type": "str", + "required": False, + "placeholder": "", + "list": False, + "show": False, + "multline": False, + "value": "f-string", + }, + "validate_template": { + "type": "bool", + "required": False, + "placeholder": "", + "list": False, + "show": False, + "multline": False, + "value": True, + }, + }, + "description": "Schema to represent a prompt for an LLM.", + "base_classes": ["BasePromptTemplate"], + } + + return prompt diff --git a/langflow/backend/list_endpoints.py b/langflow/backend/list_endpoints.py index 5448b3613..9d2acbdd8 100644 --- a/langflow/backend/list_endpoints.py +++ b/langflow/backend/list_endpoints.py @@ -7,6 +7,7 @@ from langchain import llms from langchain.chains.conversation import memory as memories from langchain.agents.load_tools import get_all_tool_names from langflow.backend import util +from langflow.backend import customs # build router @@ -52,10 +53,12 @@ def list_agents(): @router.get("/prompts") def list_prompts(): """List all prompt types""" - return [ + custom_prompts = customs.get_custom_prompts() + library_prompts = [ prompt.__annotations__["return"].__name__ for prompt in prompts.loading.type_to_loader_dict.values() ] + return library_prompts + list(custom_prompts.keys()) @router.get("/llms") diff --git a/langflow/backend/payload.py b/langflow/backend/payload.py index f66c8a504..0b25c10a6 100644 --- a/langflow/backend/payload.py +++ b/langflow/backend/payload.py @@ -59,12 +59,16 @@ def build_json(root, nodes, edges): # if module_type == "Tool": # pass if module_type in ["str", "bool", "int", "float", "Any"]: + # print(key) + # try: value = value["value"] + # except: + # pass elif "dict" in module_type: value = {} else: # if value['list']: - print(key) + # print(key) children = [] for c in local_nodes: module_types = [c["data"]["type"]] diff --git a/langflow/backend/signature.py b/langflow/backend/signature.py index bfa2c94f6..9d32a837a 100644 --- a/langflow/backend/signature.py +++ b/langflow/backend/signature.py @@ -10,6 +10,7 @@ from langchain.agents.load_tools import ( from langchain.chains.conversation import memory as memories from langflow.backend import util +from langflow.backend import customs # build router router = APIRouter( @@ -42,6 +43,8 @@ def get_agent(name: str): def get_prompt(name: str): """Get the signature of a prompt.""" try: + if name in customs.get_custom_prompts().keys(): + return customs.get_custom_prompts()[name] return util.build_template_from_function( name, prompts.loading.type_to_loader_dict )