🔧 chore(config): add VectorStoreToolkit to toolkits list

🐛 fix(base.py): remove deepcopy for VectorStore and VectorStoreRouter agents
🐛 fix(nodes.py): remove deepcopy for VectorStore and VectorStoreRouter agents
🔧 chore(loading.py): comment out unused code for loading toolkits
🐛 fix(toolkits/base.py): add Tool to base_classes in get_signature method
The changes to the config file add the VectorStoreToolkit to the list of toolkits. The deepcopy for VectorStore and VectorStoreRouter agents was causing issues, so it was removed from the base.py and nodes.py files. The loading.py file had some unused code for loading toolkits, so it was commented out. Finally, the base.py file had a bug where the Tool class was not being added to the base_classes list in the get_signature method, so it was added.
This commit is contained in:
Gabriel Almeida 2023-05-30 00:38:14 -03:00
commit a9c7fc0a69
5 changed files with 17 additions and 25 deletions

View file

@ -74,6 +74,7 @@ toolkits:
- JsonToolkit - JsonToolkit
- VectorStoreInfo - VectorStoreInfo
- VectorStoreRouterToolkit - VectorStoreRouterToolkit
- VectorStoreToolkit
tools: tools:
- Search - Search
- PAL-MATH - PAL-MATH

View file

@ -212,19 +212,7 @@ class Node:
if not self._built or force: if not self._built or force:
self._build() self._build()
#! Deepcopy is breaking for vectorstores return self._built_object
if self.base_type in [
"vectorstores",
"VectorStoreRouterAgent",
"VectorStoreAgent",
"VectorStoreInfo",
] or self.node_type in [
"VectorStoreInfo",
"VectorStoreRouterToolkit",
"SQLDatabase",
]:
return self._built_object
return deepcopy(self._built_object)
def add_edge(self, edge: "Edge") -> None: def add_edge(self, edge: "Edge") -> None:
self.edges.append(edge) self.edges.append(edge)

View file

@ -14,7 +14,7 @@ class AgentNode(Node):
def _set_tools_and_chains(self) -> None: def _set_tools_and_chains(self) -> None:
for edge in self.edges: for edge in self.edges:
source_node = edge.source source_node = edge.source
if isinstance(source_node, ToolNode): if isinstance(source_node, (ToolNode, ToolkitNode)):
self.tools.append(source_node) self.tools.append(source_node)
elif isinstance(source_node, ChainNode): elif isinstance(source_node, ChainNode):
self.chains.append(source_node) self.chains.append(source_node)
@ -32,9 +32,6 @@ class AgentNode(Node):
self._build() self._build()
#! Cannot deepcopy VectorStore, VectorStoreRouter, or SQL agents
if self.node_type in ["VectorStoreAgent", "VectorStoreRouterAgent", "SQLAgent"]:
return self._built_object
return self._built_object return self._built_object

View file

@ -101,8 +101,11 @@ def instantiate_tool(node_type, class_object, params):
def instantiate_toolkit(node_type, class_object, params): def instantiate_toolkit(node_type, class_object, params):
loaded_toolkit = class_object(**params) loaded_toolkit = class_object(**params)
if toolkits_creator.has_create_function(node_type): # Commenting this out for now to use toolkits as normal tools
return load_toolkits_executor(node_type, loaded_toolkit, params) # if toolkits_creator.has_create_function(node_type):
# return load_toolkits_executor(node_type, loaded_toolkit, params)
if isinstance(loaded_toolkit, BaseToolkit):
return loaded_toolkit.get_tools()
return loaded_toolkit return loaded_toolkit

View file

@ -42,24 +42,27 @@ class ToolkitCreator(LangChainTypeCreator):
def get_signature(self, name: str) -> Optional[Dict]: def get_signature(self, name: str) -> Optional[Dict]:
try: try:
return build_template_from_class(name, self.type_to_loader_dict) template = build_template_from_class(name, self.type_to_loader_dict)
# add Tool to base_classes
if template:
template["base_classes"].append("Tool")
return template
except ValueError as exc: except ValueError as exc:
raise ValueError("Prompt not found") from exc raise ValueError("Toolkit not found") from exc
except AttributeError as exc: except AttributeError as exc:
logger.error(f"Prompt {name} not loaded: {exc}") logger.error(f"Toolkit {name} not loaded: {exc}")
return None return None
def to_list(self) -> List[str]: def to_list(self) -> List[str]:
return list(self.type_to_loader_dict.keys()) return list(self.type_to_loader_dict.keys())
def get_create_function(self, name: str) -> Callable: def get_create_function(self, name: str) -> Callable:
if loader_name := self.create_functions.get(name, None): if loader_name := self.create_functions.get(name):
# import loader
return import_module( return import_module(
f"from langchain.agents.agent_toolkits import {loader_name[0]}" f"from langchain.agents.agent_toolkits import {loader_name[0]}"
) )
else: else:
raise ValueError("Loader not found") raise ValueError("Toolkit not found")
def has_create_function(self, name: str) -> bool: def has_create_function(self, name: str) -> bool:
# check if the function list is not empty # check if the function list is not empty