feat: implemented load
This commit is contained in:
parent
38ce27352d
commit
784579e06a
7 changed files with 157 additions and 23 deletions
16
poetry.lock
generated
16
poetry.lock
generated
|
|
@ -268,6 +268,17 @@ category = "main"
|
||||||
optional = false
|
optional = false
|
||||||
python-versions = ">=3.7"
|
python-versions = ">=3.7"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "google-search-results"
|
||||||
|
version = "2.4.1"
|
||||||
|
description = "Scrape and search localized results from Google, Bing, Baidu, Yahoo, Yandex, Ebay, Homedepot, youtube at scale using SerpApi.com"
|
||||||
|
category = "main"
|
||||||
|
optional = false
|
||||||
|
python-versions = ">=3.5"
|
||||||
|
|
||||||
|
[package.dependencies]
|
||||||
|
requests = "*"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "greenlet"
|
name = "greenlet"
|
||||||
version = "2.0.2"
|
version = "2.0.2"
|
||||||
|
|
@ -933,7 +944,7 @@ multidict = ">=4.0"
|
||||||
[metadata]
|
[metadata]
|
||||||
lock-version = "1.1"
|
lock-version = "1.1"
|
||||||
python-versions = "^3.10"
|
python-versions = "^3.10"
|
||||||
content-hash = "ffc4b17c403dab7f934d2c026b83757cc40066a79d64ab27842908b40b0bedeb"
|
content-hash = "a8de7c079509305c680e1408fc12648c5ccda4963d55cc6d5a31758dad702688"
|
||||||
|
|
||||||
[metadata.files]
|
[metadata.files]
|
||||||
aiohttp = [
|
aiohttp = [
|
||||||
|
|
@ -1282,6 +1293,9 @@ frozenlist = [
|
||||||
{file = "frozenlist-1.3.3-cp39-cp39-win_amd64.whl", hash = "sha256:cfe33efc9cb900a4c46f91a5ceba26d6df370ffddd9ca386eb1d4f0ad97b9ea9"},
|
{file = "frozenlist-1.3.3-cp39-cp39-win_amd64.whl", hash = "sha256:cfe33efc9cb900a4c46f91a5ceba26d6df370ffddd9ca386eb1d4f0ad97b9ea9"},
|
||||||
{file = "frozenlist-1.3.3.tar.gz", hash = "sha256:58bcc55721e8a90b88332d6cd441261ebb22342e238296bb330968952fbb3a6a"},
|
{file = "frozenlist-1.3.3.tar.gz", hash = "sha256:58bcc55721e8a90b88332d6cd441261ebb22342e238296bb330968952fbb3a6a"},
|
||||||
]
|
]
|
||||||
|
google-search-results = [
|
||||||
|
{file = "google_search_results-2.4.1.tar.gz", hash = "sha256:021746fc21c0b0786e61a2d103d93a08c5c84e204d3f93cd4d589e0e117614a7"},
|
||||||
|
]
|
||||||
greenlet = [
|
greenlet = [
|
||||||
{file = "greenlet-2.0.2-cp27-cp27m-macosx_10_14_x86_64.whl", hash = "sha256:bdfea8c661e80d3c1c99ad7c3ff74e6e87184895bbaca6ee8cc61209f8b9b85d"},
|
{file = "greenlet-2.0.2-cp27-cp27m-macosx_10_14_x86_64.whl", hash = "sha256:bdfea8c661e80d3c1c99ad7c3ff74e6e87184895bbaca6ee8cc61209f8b9b85d"},
|
||||||
{file = "greenlet-2.0.2-cp27-cp27m-manylinux2010_x86_64.whl", hash = "sha256:9d14b83fab60d5e8abe587d51c75b252bcc21683f24699ada8fb275d7712f5a9"},
|
{file = "greenlet-2.0.2-cp27-cp27m-manylinux2010_x86_64.whl", hash = "sha256:9d14b83fab60d5e8abe587d51c75b252bcc21683f24699ada8fb275d7712f5a9"},
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ fastapi = "^0.91.0"
|
||||||
uvicorn = "^0.20.0"
|
uvicorn = "^0.20.0"
|
||||||
beautifulsoup4 = "^4.11.2"
|
beautifulsoup4 = "^4.11.2"
|
||||||
langchain = {path = "../langchain", develop = true}
|
langchain = {path = "../langchain", develop = true}
|
||||||
|
google-search-results = "^2.4.1"
|
||||||
|
|
||||||
|
|
||||||
[tool.poetry.group.dev.dependencies]
|
[tool.poetry.group.dev.dependencies]
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,10 @@
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
import signature
|
import signature
|
||||||
import list_endpoints
|
import list_endpoints
|
||||||
|
import payload
|
||||||
|
from langchain.agents.loading import load_agent_executor_from_config
|
||||||
|
from langchain.prompts.loading import load_prompt_from_config
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
@ -21,22 +25,23 @@ def get_type_list():
|
||||||
def get_all():
|
def get_all():
|
||||||
return {
|
return {
|
||||||
"chains": {
|
"chains": {
|
||||||
chain: signature.chain(chain) for chain in list_endpoints.list_chains()
|
chain: signature.get_chain(chain) for chain in list_endpoints.list_chains()
|
||||||
},
|
},
|
||||||
"agents": {
|
"agents": {
|
||||||
agent: signature.agent(agent) for agent in list_endpoints.list_agents()
|
agent: signature.get_agent(agent) for agent in list_endpoints.list_agents()
|
||||||
},
|
},
|
||||||
"prompts": {
|
"prompts": {
|
||||||
prompt: signature.prompt(prompt) for prompt in list_endpoints.list_prompts()
|
prompt: signature.get_prompt(prompt)
|
||||||
|
for prompt in list_endpoints.list_prompts()
|
||||||
},
|
},
|
||||||
"llms": {llm: signature.llm(llm) for llm in list_endpoints.list_llms()},
|
"llms": {llm: signature.get_llm(llm) for llm in list_endpoints.list_llms()},
|
||||||
# "utilities": {
|
# "utilities": {
|
||||||
# "template": {
|
# "template": {
|
||||||
# # utility: templates.utility(utility) for utility in list.list_utilities()
|
# # utility: templates.utility(utility) for utility in list.list_utilities()
|
||||||
# }
|
# }
|
||||||
# },
|
# },
|
||||||
"memories": {
|
"memories": {
|
||||||
memory: signature.memory(memory)
|
memory: signature.get_memory(memory)
|
||||||
for memory in list_endpoints.list_memories()
|
for memory in list_endpoints.list_memories()
|
||||||
},
|
},
|
||||||
# "document_loaders": {
|
# "document_loaders": {
|
||||||
|
|
@ -51,16 +56,38 @@ def get_all():
|
||||||
# tool: {"template": signature.tool(tool), **values}
|
# tool: {"template": signature.tool(tool), **values}
|
||||||
# for tool, values in tools.items()
|
# for tool, values in tools.items()
|
||||||
# },
|
# },
|
||||||
"tools": {tool: signature.tool(tool) for tool in list_endpoints.list_tools()},
|
"tools": {
|
||||||
|
tool: signature.get_tool(tool) for tool in list_endpoints.list_tools()
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@router.post("/predict")
|
@router.post("/predict")
|
||||||
def get_load(data: dict[str, str]):
|
def get_load(data: dict[str, Any]):
|
||||||
a = get_type_list()
|
type_list = get_type_list()
|
||||||
|
|
||||||
|
# Add input variables
|
||||||
|
data = payload.extract_input_variables(data)
|
||||||
|
|
||||||
|
# Nodes, edges and root node
|
||||||
|
message = data["message"]
|
||||||
|
nodes = data["nodes"]
|
||||||
|
edges = data["edges"]
|
||||||
|
root = payload.get_root_node(data)
|
||||||
|
|
||||||
|
extracted_json = payload.build_json(root, nodes, edges)
|
||||||
|
|
||||||
# Build json
|
# Build json
|
||||||
|
if extracted_json["_type"] in type_list["agents"]:
|
||||||
|
loaded = load_agent_executor_from_config(extracted_json)
|
||||||
|
|
||||||
|
return loaded.run(message)
|
||||||
|
|
||||||
|
elif extracted_json["_type"] in type_list["prompts"]:
|
||||||
|
loaded = load_prompt_from_config(extracted_json)
|
||||||
|
print(loaded.format(product=''))
|
||||||
|
return extracted_json
|
||||||
|
|
||||||
# if type in a["prompts"]:
|
# if type in a["prompts"]:
|
||||||
|
|
||||||
return a
|
# return a
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ from langchain import prompts
|
||||||
from langchain import llms
|
from langchain import llms
|
||||||
from langchain.chains.conversation import memory as memories
|
from langchain.chains.conversation import memory as memories
|
||||||
from langchain.agents.load_tools import get_all_tool_names
|
from langchain.agents.load_tools import get_all_tool_names
|
||||||
|
import util
|
||||||
|
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
|
|
@ -97,4 +98,7 @@ def list_memories():
|
||||||
def list_tools():
|
def list_tools():
|
||||||
"""List all load tools"""
|
"""List all load tools"""
|
||||||
|
|
||||||
return get_all_tool_names()
|
return [
|
||||||
|
util.get_tool_params(util.get_tools_dict(tool))["name"]
|
||||||
|
for tool in get_all_tool_names()
|
||||||
|
]
|
||||||
|
|
|
||||||
80
src/payload.py
Normal file
80
src/payload.py
Normal file
|
|
@ -0,0 +1,80 @@
|
||||||
|
import re
|
||||||
|
|
||||||
|
|
||||||
|
def extract_input_variables(data):
|
||||||
|
"""
|
||||||
|
Extracts input variables from the template and adds them to the input_variables field.
|
||||||
|
"""
|
||||||
|
for node in data["nodes"]:
|
||||||
|
try:
|
||||||
|
if "input_variables" in node["data"]["node"]["template"]:
|
||||||
|
if node["data"]["node"]["template"]["_type"] == "prompt":
|
||||||
|
variables = re.findall(
|
||||||
|
r"\{(.*?)\}",
|
||||||
|
node["data"]["node"]["template"]["template"]["value"],
|
||||||
|
)
|
||||||
|
elif node["data"]["node"]["template"]["_type"] == "few_shot":
|
||||||
|
variables = re.findall(
|
||||||
|
r"\{(.*?)\}",
|
||||||
|
node["data"]["node"]["template"]["prefix"]["value"]
|
||||||
|
+ node["data"]["node"]["template"]["suffix"]["value"],
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
variables = []
|
||||||
|
node["data"]["node"]["template"]["input_variables"]["value"] = variables
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def get_root_node(data):
|
||||||
|
"""
|
||||||
|
Returns the root node of the template.
|
||||||
|
"""
|
||||||
|
root = None
|
||||||
|
incoming_edges = {edge["source"] for edge in data["edges"]}
|
||||||
|
for node in data["nodes"]:
|
||||||
|
if node["id"] not in incoming_edges:
|
||||||
|
root = node
|
||||||
|
break
|
||||||
|
return root
|
||||||
|
|
||||||
|
|
||||||
|
def build_json(root, nodes, edges):
|
||||||
|
edge_ids = [edge["source"] for edge in edges if edge["target"] == root["id"]]
|
||||||
|
local_nodes = [node for node in nodes if node["id"] in edge_ids]
|
||||||
|
|
||||||
|
if "node" not in root["data"]:
|
||||||
|
return build_json(local_nodes[0], nodes, edges)
|
||||||
|
|
||||||
|
final_dict = root["data"]["node"]["template"].copy()
|
||||||
|
|
||||||
|
for key, value in final_dict.items():
|
||||||
|
if key == "_type":
|
||||||
|
continue
|
||||||
|
|
||||||
|
module_type = value["type"]
|
||||||
|
if module_type == "Tool":
|
||||||
|
pass
|
||||||
|
if module_type in ["str", "bool", "int", "float"]:
|
||||||
|
value = value["value"]
|
||||||
|
elif "dict" in module_type:
|
||||||
|
value = {}
|
||||||
|
else:
|
||||||
|
# if value['list']:
|
||||||
|
children = [
|
||||||
|
c
|
||||||
|
for c in local_nodes
|
||||||
|
if module_type
|
||||||
|
in [c["data"]["type"]] + c["data"]["node"]["base_classes"]
|
||||||
|
]
|
||||||
|
# else:
|
||||||
|
# children = next((c for c in local_nodes if type in [c['data']['type']] + c['data']['node']['base_classes']), None)
|
||||||
|
if value["required"] and not children:
|
||||||
|
raise ValueError(f"No child with type {module_type} found")
|
||||||
|
values = [
|
||||||
|
build_json(child, nodes, edges) for child in children
|
||||||
|
] # if children else None
|
||||||
|
value = list(values) if value["list"] else next(iter(values), None)
|
||||||
|
final_dict[key] = value
|
||||||
|
return final_dict
|
||||||
|
|
@ -172,8 +172,14 @@ def get_memory(name: str):
|
||||||
@router.get("/tool")
|
@router.get("/tool")
|
||||||
def get_tool(name: str):
|
def get_tool(name: str):
|
||||||
"""Get the signature of a tool."""
|
"""Get the signature of a tool."""
|
||||||
|
|
||||||
|
all_tools = {
|
||||||
|
util.get_tool_params(util.get_tools_dict(tool))["name"]: tool
|
||||||
|
for tool in get_all_tool_names()
|
||||||
|
}
|
||||||
|
|
||||||
# Raise error if name is not in tools
|
# Raise error if name is not in tools
|
||||||
if name not in get_all_tool_names():
|
if name not in all_tools.keys():
|
||||||
raise HTTPException(status_code=404, detail=f"Tool {name} not found.")
|
raise HTTPException(status_code=404, detail=f"Tool {name} not found.")
|
||||||
|
|
||||||
type_dict = {
|
type_dict = {
|
||||||
|
|
@ -188,26 +194,28 @@ def get_tool(name: str):
|
||||||
"llm": {"type": "BaseLLM", "required": True, "list": False, "show": True},
|
"llm": {"type": "BaseLLM", "required": True, "list": False, "show": True},
|
||||||
}
|
}
|
||||||
|
|
||||||
if name in _BASE_TOOLS:
|
tool_type = all_tools[name]
|
||||||
|
|
||||||
|
if tool_type in _BASE_TOOLS:
|
||||||
params = []
|
params = []
|
||||||
elif name in _LLM_TOOLS:
|
elif tool_type in _LLM_TOOLS:
|
||||||
params = ["llm"]
|
params = ["llm"]
|
||||||
elif name in _EXTRA_LLM_TOOLS:
|
elif tool_type in _EXTRA_LLM_TOOLS:
|
||||||
_, extra_keys = _EXTRA_LLM_TOOLS[name]
|
_, extra_keys = _EXTRA_LLM_TOOLS[tool_type]
|
||||||
params = ["llm"] + extra_keys
|
params = ["llm"] + extra_keys
|
||||||
elif name in _EXTRA_OPTIONAL_TOOLS:
|
elif tool_type in _EXTRA_OPTIONAL_TOOLS:
|
||||||
_, extra_keys = _EXTRA_OPTIONAL_TOOLS[name]
|
_, extra_keys = _EXTRA_OPTIONAL_TOOLS[tool_type]
|
||||||
params = extra_keys
|
params = extra_keys
|
||||||
|
|
||||||
template = {
|
template = {
|
||||||
param: (type_dict[param] if param == "llm" else type_dict["str"])
|
param: (type_dict[param] if param == "llm" else type_dict["str"])
|
||||||
for param in params
|
for param in params
|
||||||
}
|
}
|
||||||
template["_type"] = name
|
template["_type"] = tool_type
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"template": template,
|
"template": template,
|
||||||
**util.get_tool_params(util.get_tools_dict(name)),
|
**util.get_tool_params(util.get_tools_dict(tool_type)),
|
||||||
"base_classes": ["Tool"],
|
"base_classes": ["Tool"],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -135,6 +135,8 @@ def format_dict(d):
|
||||||
|
|
||||||
# Process remaining keys
|
# Process remaining keys
|
||||||
for key, value in d.items():
|
for key, value in d.items():
|
||||||
|
if key == "examples":
|
||||||
|
pass
|
||||||
if key == "_type":
|
if key == "_type":
|
||||||
continue
|
continue
|
||||||
_type = value["type"]
|
_type = value["type"]
|
||||||
|
|
@ -176,6 +178,4 @@ def format_dict(d):
|
||||||
value.pop("default")
|
value.pop("default")
|
||||||
|
|
||||||
# Filter out keys that should not be shown
|
# Filter out keys that should not be shown
|
||||||
return (
|
return d
|
||||||
d # {key: value for key, value in d.items() if key == "_type" or value["show"]}
|
|
||||||
)
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue