diff --git a/server/routers/base_router.py b/server/routers/base_router.py
index cbf8d1fe..8ee6f328 100644
--- a/server/routers/base_router.py
+++ b/server/routers/base_router.py
@@ -28,7 +28,7 @@ async def update_config(key = Body(...), value = Body(...)):
@base.post("/restart")
async def restart():
knowledge_base.restart()
- graph_base.restart()
+ graph_base.start()
retriever.restart()
return {"message": "Restarted!"}
diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py
index 55959402..4c089d7d 100644
--- a/server/routers/chat_router.py
+++ b/server/routers/chat_router.py
@@ -12,6 +12,7 @@ from src.core import HistoryManager
from src.agents import agent_manager
from src.models import select_model
from src.utils.logging_config import logger
+from src.agents.tools_factory import get_all_tools
chat = APIRouter(prefix="/chat")
@@ -201,3 +202,8 @@ async def get_chat_models(model_provider: str):
"""获取指定模型提供商的模型列表"""
model = select_model(model_provider=model_provider)
return {"models": model.get_models()}
+
+@chat.get("/tools")
+async def get_tools():
+ """获取所有工具"""
+ return {"tools": list(get_all_tools().keys())}
diff --git a/src/agents/chatbot/configuration.py b/src/agents/chatbot/configuration.py
index fbdb68bc..20eac6e0 100644
--- a/src/agents/chatbot/configuration.py
+++ b/src/agents/chatbot/configuration.py
@@ -34,3 +34,12 @@ class ChatbotConfiguration(Configuration):
"description": "智能体的驱动模型"
},
)
+
+ tools: list[str] = field(
+ default_factory=list,
+ metadata={
+ "name": "工具",
+ "configurable": False,
+ "description": "工具列表"
+ },
+ )
diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py
index 1d4162ce..a80b6d38 100644
--- a/src/agents/chatbot/graph.py
+++ b/src/agents/chatbot/graph.py
@@ -13,31 +13,32 @@ from src.utils import logger
from src.agents.registry import State, BaseAgent
from src.agents.utils import load_chat_model, get_cur_time_with_utc
from src.agents.chatbot.configuration import ChatbotConfiguration
-from src.agents.tools_factory import _TOOLS_REGISTRY
+from src.agents.tools_factory import get_all_tools
class ChatbotAgent(BaseAgent):
name = "chatbot"
- description = "A chatbot that can answer questions and help with tasks."
+ description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。"
requirements = ["TAVILY_API_KEY", "ZHIPUAI_API_KEY"]
- all_tools = ["TavilySearchResults", "multiply", "add", "subtract", "divide"]
config_schema = ChatbotConfiguration
def __init__(self, **kwargs):
super().__init__(**kwargs)
- def _get_tools(self, config_schema: RunnableConfig):
- """根据配置获取工具,如果配置为空,则使用所有工具,如果配置为列表,则使用列表中的工具,
- 如果配置为其他类型,则抛出错误"""
- conf_tools = config_schema.get("tools")
- if conf_tools == None:
- tool_names = self.all_tools
- elif isinstance(conf_tools, list):
- tool_names = [tool for tool in self.all_tools if tool in conf_tools]
+ def _get_tools(self, tools: list[str]):
+ """根据配置获取工具。
+ 默认不使用任何工具。
+ 如果配置为列表,则使用列表中的工具。
+ """
+ platform_tools = get_all_tools()
+ if tools is None or not isinstance(tools, list) or len(tools) == 0:
+ # 默认不使用任何工具
+ logger.info("未配置工具或配置为空,不使用任何工具")
+ return []
else:
- raise ValueError(f"tools 配置错误: {conf_tools}")
-
- logger.info(f"Tools: {tool_names}")
- return [_TOOLS_REGISTRY[tool] for tool in tool_names]
+ # 使用配置中指定的工具
+ tool_names = [tool for tool in platform_tools.keys() if tool in tools]
+ logger.info(f"使用工具: {tool_names}")
+ return [platform_tools[tool] for tool in tool_names]
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
"""调用 llm 模型"""
@@ -46,7 +47,7 @@ class ChatbotAgent(BaseAgent):
system_prompt = f"{conf.system_prompt} Now is {get_cur_time_with_utc()}"
model = load_chat_model(conf.model)
- model_with_tools = model.bind_tools(self._get_tools(config_schema))
+ model_with_tools = model.bind_tools(self._get_tools(conf.tools))
logger.info(f"llm_call with config: {conf}, {conf.model}")
res = model_with_tools.invoke(
@@ -56,9 +57,10 @@ class ChatbotAgent(BaseAgent):
def get_graph(self, config_schema: RunnableConfig = None, **kwargs):
"""构建图"""
+ conf = self.config_schema.from_runnable_config(config_schema)
workflow = StateGraph(State, config_schema=self.config_schema)
workflow.add_node("chatbot", self.llm_call)
- workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
+ workflow.add_node("tools", ToolNode(tools=self._get_tools(conf.tools)))
workflow.add_edge(START, "chatbot")
workflow.add_conditional_edges(
"chatbot",
diff --git a/src/agents/react/workflows.py b/src/agents/react/workflows.py
index 1337d133..ae9db167 100644
--- a/src/agents/react/workflows.py
+++ b/src/agents/react/workflows.py
@@ -2,6 +2,8 @@ import os
from langchain_openai import ChatOpenAI
+from src import graph_base
+
model = ChatOpenAI(model="glm-4-plus",
api_key=os.getenv("ZHIPUAI_API_KEY"),
base_url="https://open.bigmodel.cn/api/paas/v4/",
@@ -10,26 +12,14 @@ model = ChatOpenAI(model="glm-4-plus",
# For this tutorial we will use custom tool that returns pre-defined values for weather in two cities (NYC & SF)
-from typing import Literal
+from typing import Literal, Annotated
-from langchain_core.tools import tool
+from langchain_core.tools import tool, StructuredTool
-@tool
-def get_weather(city: Literal["nyc", "sf"]):
- """Use this to get weather information."""
- if city == "nyc":
- return "It might be cloudy in nyc"
- elif city == "sf":
- return "It's always sunny in sf"
- else:
- raise AssertionError("Unknown city")
+tools = []
-tools = [get_weather]
-
-
-# Define the graph
from langgraph.prebuilt import create_react_agent
diff --git a/src/agents/tools_factory.py b/src/agents/tools_factory.py
index 61369af4..21fa07cc 100644
--- a/src/agents/tools_factory.py
+++ b/src/agents/tools_factory.py
@@ -1,13 +1,15 @@
import json
import re
import os
-from typing import Any, Callable, Optional, Type, Union
+from typing import Any, Callable, Optional, Type, Union, Annotated
from pydantic import BaseModel, Field
-from langchain_core.tools import tool, BaseTool
+from langchain_core.tools import tool, BaseTool, StructuredTool
from langchain_community.tools.tavily_search import TavilySearchResults
+from src import graph_base, knowledge_base
+
# refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py
def regist_tool(
*args: Any,
@@ -63,6 +65,24 @@ def regist_tool(
return t
+class KnowledgeRetrieverModel(BaseModel):
+ query: str = Field(description="The query to get knowledge graph.")
+
+
+
+def get_all_tools():
+ """获取所有工具"""
+ tools = _TOOLS_REGISTRY.copy()
+ for db_Id, retrieve_info in knowledge_base.get_retrievers().items():
+ name = f"retrieve_{retrieve_info['name']}"
+ tools[name] = StructuredTool.from_function(
+ retrieve_info["retriever"],
+ name=name,
+ description=retrieve_info["description"],
+ args_schema=KnowledgeRetrieverModel)
+
+ return tools
+
class BaseToolOutput:
"""
LLM 要求 Tool 的输出为 str,但 Tool 用在别处时希望它正常返回结构化数据。
@@ -92,33 +112,30 @@ class BaseToolOutput:
else:
return str(self.data)
-
+@tool
+def calculator(a: float, b: float, operation: str) -> float:
+ """Calculate two numbers."""
+ if operation == "add":
+ return a + b
+ elif operation == "subtract":
+ return a - b
+ elif operation == "multiply":
+ return a * b
+ elif operation == "divide":
+ return a / b
+ else:
+ raise ValueError(f"Invalid operation: {operation}, only support add, subtract, multiply, divide")
@tool
-def multiply(first_int: int, second_int: int) -> int:
- """Multiply two integers together."""
- return first_int * second_int
+def get_knowledge_graph(query: Annotated[str, "The query to get knowledge graph."]):
+ """Use this to get knowledge graph."""
+ return graph_base.query_node(query, hops=2)
-@tool
-def add(first_int: int, second_int: int) -> int:
- """Add two integers together."""
- return first_int + second_int
-@tool
-def subtract(first_int: int, second_int: int) -> int:
- """Subtract two integers."""
- return first_int - second_int
-
-@tool
-def divide(first_int: int, second_int: int) -> int:
- """Divide two integers."""
- return first_int / second_int
_TOOLS_REGISTRY = {
- "multiply": multiply,
- "add": add,
- "subtract": subtract,
- "divide": divide,
+ "calculator": calculator,
"TavilySearchResults": TavilySearchResults(max_results=10),
+ "get_knowledge_graph": get_knowledge_graph,
}
diff --git a/src/core/graphbase.py b/src/core/graphbase.py
index 193b99a7..6f0af518 100644
--- a/src/core/graphbase.py
+++ b/src/core/graphbase.py
@@ -209,6 +209,9 @@ class GraphDatabase:
def query_node(self, entity_name, hops=2, **kwargs):
# TODO 添加判断节点数量为 0 停止检索
+ # 判断是否启动
+ if not self.is_running():
+ raise Exception("图数据库未启动")
logger.debug(f"Query graph node {entity_name} with {hops=}")
if kwargs.get("exact_match"):
@@ -266,7 +269,7 @@ class GraphDatabase:
def query(tx, entity_name, hops):
result = tx.run(f"""
MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m)
- RETURN n, r, m
+ RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""", entity_name=entity_name)
return result.values()
@@ -279,7 +282,7 @@ class GraphDatabase:
def query(tx, hops):
result = tx.run(f"""
MATCH (n)-[r*1..{hops}]->(m)
- RETURN n, r, m
+ RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""")
return result.values()
@@ -292,7 +295,7 @@ class GraphDatabase:
def query(tx, relationship_type, hops):
result = tx.run(f"""
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
- RETURN n, r, m
+ RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""")
return result.values()
@@ -307,7 +310,7 @@ class GraphDatabase:
MATCH (n:Entity)
WHERE n.name CONTAINS $keyword
MATCH (n)-[r*1..{hops}]->(m)
- RETURN n, r, m
+ RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""", keyword=keyword)
return result.values()
@@ -321,7 +324,7 @@ class GraphDatabase:
result = tx.run(f"""
MATCH (n {{name: $node_name}})
OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
- RETURN n, r, m
+ RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""", node_name=node_name)
return result.values()
diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py
index 2c5709ee..92e70594 100644
--- a/src/core/knowledgebase.py
+++ b/src/core/knowledgebase.py
@@ -344,7 +344,7 @@ class KnowledgeBase:
"all_results": all_db_result,
}
- def get_retriever(self, db_id):
+ def get_retriever_by_db_id(self, db_id):
retriever_params = {
"distance_threshold": self.default_distance_threshold,
"rerank_threshold": self.default_rerank_threshold,
@@ -358,6 +358,16 @@ class KnowledgeBase:
return retriever
+ def get_retrievers(self):
+ retrievers = {}
+ for db in self.db_manager.get_all_databases():
+ retrievers[db["db_id"]] = {
+ "name": db["name"],
+ "description": db["description"],
+ "retriever": self.get_retriever_by_db_id(db["db_id"]),
+ }
+ return retrievers
+
################################
#* Below is the code for milvus #
################################
diff --git a/web/src/components/AgentChatComponent.vue b/web/src/components/AgentChatComponent.vue
index d28caca3..d1e072cc 100644
--- a/web/src/components/AgentChatComponent.vue
+++ b/web/src/components/AgentChatComponent.vue
@@ -360,6 +360,7 @@ const sendMessageWithText = async (text) => {
body: JSON.stringify(requestData)
});
+ // console.log("requestData", requestData);
if (!response.ok) {
throw new Error('请求失败');
}
@@ -554,7 +555,7 @@ const handleFinished = async (data) => {
const handleMessageById = async (data) => {
const msgId = data.msg.id;
const msgType = data.msg.type;
- console.log("data", data);
+ // console.log("data", data);
// 查找现有消息
const existingMsgIndex = messageMap.value.get(msgId);
diff --git a/web/src/views/AgentView.vue b/web/src/views/AgentView.vue
index 9db3c27f..2ca517d0 100644
--- a/web/src/views/AgentView.vue
+++ b/web/src/views/AgentView.vue
@@ -163,6 +163,23 @@
+
+ 选择要启用的工具