From 380e7e38d60d0f008bbc8306f8a143cbde565deb Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 11 Apr 2025 11:54:45 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E7=9F=A5=E8=AF=86=E5=BA=93?= =?UTF-8?q?=E6=9F=A5=E8=AF=A2=E4=BB=A5=E5=8F=8A=E5=B7=A5=E5=85=B7=E5=AE=9A?= =?UTF-8?q?=E4=B9=89=E6=96=B9=E6=B3=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agents/tools_factory.py | 10 ++++++++-- src/core/knowledgebase.py | 29 ++++++++++++++++++++--------- web/src/views/AgentView.vue | 2 +- web/src/views/DataBaseView.vue | 17 +++++++++-------- 4 files changed, 38 insertions(+), 20 deletions(-) diff --git a/src/agents/tools_factory.py b/src/agents/tools_factory.py index a7e9eef8..3f22eab4 100644 --- a/src/agents/tools_factory.py +++ b/src/agents/tools_factory.py @@ -66,19 +66,25 @@ def regist_tool( class KnowledgeRetrieverModel(BaseModel): - query: str = Field(description="The query to get knowledge graph.") + query: str = Field(description="查询的关键词,查询的时候,应该尽量以关键词的形式进行查询,不要使用复杂的句子。\n.") 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']}" + description = ( + f"使用 {retrieve_info['name']} 知识库进行检索。\n" + f"下面是这个知识库的描述:\n{retrieve_info['description']}" + ) tools[name] = StructuredTool.from_function( retrieve_info["retriever"], name=name, - description=retrieve_info["description"], + description=description, args_schema=KnowledgeRetrieverModel) return tools diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 92e70594..620260e6 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -11,7 +11,6 @@ from src.utils import logger, hashstr from src.core.indexing import chunk, read_text from src.core.kb_db_manager import kb_db_manager - class KnowledgeBase: def __init__(self) -> None: @@ -250,7 +249,7 @@ class KnowledgeBase: def add_files(self, db_id, files, params=None): db = self.get_kb_by_id(db_id) - if db["embed_model"] != self.embed_model.embed_model_fullname: + if not self.check_embed_model(db_id): logger.error(f"Embed model not match, {db['embed_model']} != {self.embed_model.embed_model_fullname}") return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"} @@ -312,7 +311,6 @@ class KnowledgeBase: ################################### def query(self, query, db_id, **kwargs): - db = self.get_kb_by_id(db_id) distance_threshold = kwargs.get("distance_threshold", self.default_distance_threshold) rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold) @@ -361,11 +359,19 @@ class KnowledgeBase: 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"]), - } + if self.check_embed_model(db["db_id"]): + retrievers[db["db_id"]] = { + "name": db["name"], + "description": db["description"], + "retriever": self.get_retriever_by_db_id(db["db_id"]), + "embed_model": db["embed_model"], + } + else: + logger.warning(( + f"无法将知识库 {db['name']} 转换为 Tools, 因为向量模型不匹配," + f"当前向量模型: {self.embed_model.embed_model_fullname}," + f"知识库向量模型: {db['embed_model']}。" + )) return retrievers ################################ @@ -474,4 +480,9 @@ class KnowledgeBase: def search_by_id(self, collection_name, id, output_fields=["id", "text"]): res = self.client.get(collection_name, id, output_fields=output_fields) - return res \ No newline at end of file + return res + + def check_embed_model(self, db_id): + db = self.db_manager.get_database_by_id(db_id) + return db["embed_model"] == self.embed_model.embed_model_fullname + diff --git a/web/src/views/AgentView.vue b/web/src/views/AgentView.vue index 99c06ba5..7a2fcc43 100644 --- a/web/src/views/AgentView.vue +++ b/web/src/views/AgentView.vue @@ -165,7 +165,7 @@ -

选择要启用的工具

+

选择要启用的工具(注:retrieve 工具仅展现了与当前向量模型匹配的知识库,详情请查看 docker 日志。)

diff --git a/web/src/views/DataBaseView.vue b/web/src/views/DataBaseView.vue index 0e83c2b6..a47328b5 100644 --- a/web/src/views/DataBaseView.vue +++ b/web/src/views/DataBaseView.vue @@ -5,18 +5,19 @@ description="知识型数据库,主要是非结构化的文本组成,使用向量检索使用。如果出现问题,可以检查 saves/data/database.json 查看配置。" > - -

数据库名称*

- -

数据库描述

+ +

知识库名称*

+ +

知识库描述

+

在智能体流程中,这里的描述会作为工具的描述。智能体会根据知识库的标题和描述来选择合适的工具。所以这里描述的越详细,智能体越容易选择到合适的工具。