优化知识库查询以及工具定义方法

This commit is contained in:
Wenjie Zhang 2025-04-11 11:54:45 +08:00
parent f11f719b4e
commit 380e7e38d6
4 changed files with 38 additions and 20 deletions

View File

@ -66,19 +66,25 @@ def regist_tool(
class KnowledgeRetrieverModel(BaseModel): class KnowledgeRetrieverModel(BaseModel):
query: str = Field(description="The query to get knowledge graph.") query: str = Field(description="查询的关键词,查询的时候,应该尽量以关键词的形式进行查询,不要使用复杂的句子。\n.")
def get_all_tools(): def get_all_tools():
"""获取所有工具""" """获取所有工具"""
tools = _TOOLS_REGISTRY.copy() tools = _TOOLS_REGISTRY.copy()
# 获取所有知识库
for db_Id, retrieve_info in knowledge_base.get_retrievers().items(): for db_Id, retrieve_info in knowledge_base.get_retrievers().items():
name = f"retrieve_{retrieve_info['name']}" name = f"retrieve_{retrieve_info['name']}"
description = (
f"使用 {retrieve_info['name']} 知识库进行检索。\n"
f"下面是这个知识库的描述:\n{retrieve_info['description']}"
)
tools[name] = StructuredTool.from_function( tools[name] = StructuredTool.from_function(
retrieve_info["retriever"], retrieve_info["retriever"],
name=name, name=name,
description=retrieve_info["description"], description=description,
args_schema=KnowledgeRetrieverModel) args_schema=KnowledgeRetrieverModel)
return tools return tools

View File

@ -11,7 +11,6 @@ from src.utils import logger, hashstr
from src.core.indexing import chunk, read_text from src.core.indexing import chunk, read_text
from src.core.kb_db_manager import kb_db_manager from src.core.kb_db_manager import kb_db_manager
class KnowledgeBase: class KnowledgeBase:
def __init__(self) -> None: def __init__(self) -> None:
@ -250,7 +249,7 @@ class KnowledgeBase:
def add_files(self, db_id, files, params=None): def add_files(self, db_id, files, params=None):
db = self.get_kb_by_id(db_id) 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}") 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"} 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): 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) distance_threshold = kwargs.get("distance_threshold", self.default_distance_threshold)
rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold) rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold)
@ -361,11 +359,19 @@ class KnowledgeBase:
def get_retrievers(self): def get_retrievers(self):
retrievers = {} retrievers = {}
for db in self.db_manager.get_all_databases(): for db in self.db_manager.get_all_databases():
retrievers[db["db_id"]] = { if self.check_embed_model(db["db_id"]):
"name": db["name"], retrievers[db["db_id"]] = {
"description": db["description"], "name": db["name"],
"retriever": self.get_retriever_by_db_id(db["db_id"]), "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 return retrievers
################################ ################################
@ -474,4 +480,9 @@ class KnowledgeBase:
def search_by_id(self, collection_name, id, output_fields=["id", "text"]): def search_by_id(self, collection_name, id, output_fields=["id", "text"]):
res = self.client.get(collection_name, id, output_fields=output_fields) res = self.client.get(collection_name, id, output_fields=output_fields)
return res 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

View File

@ -165,7 +165,7 @@
<!-- 添加工具选择部分 --> <!-- 添加工具选择部分 -->
<a-form-item label="可用工具" name="tools" class="config-item"> <a-form-item label="可用工具" name="tools" class="config-item">
<p class="description">选择要启用的工具</p> <p class="description">选择要启用的工具retrieve 工具仅展现了与当前向量模型匹配的知识库详情请查看 docker 日志</p>
<a-form-item-rest> <a-form-item-rest>
<div class="tools-switches"> <div class="tools-switches">
<div v-for="tool in availableTools" :key="tool" class="tool-switch-item"> <div v-for="tool in availableTools" :key="tool" class="tool-switch-item">

View File

@ -5,18 +5,19 @@
description="知识型数据库,主要是非结构化的文本组成,使用向量检索使用。如果出现问题,可以检查 saves/data/database.json 查看配置。" description="知识型数据库,主要是非结构化的文本组成,使用向量检索使用。如果出现问题,可以检查 saves/data/database.json 查看配置。"
> >
<template #actions> <template #actions>
<a-button type="primary" @click="newDatabase.open=true">新建数据</a-button> <a-button type="primary" @click="newDatabase.open=true">新建知识</a-button>
</template> </template>
</HeaderComponent> </HeaderComponent>
<a-modal :open="newDatabase.open" title="新建数据库" @ok="createDatabase" @cancel="newDatabase.open=false"> <a-modal :open="newDatabase.open" title="新建知识库" @ok="createDatabase" @cancel="newDatabase.open=false">
<h3>数据库名称<span style="color: var(--error-color)">*</span></h3> <h3>知识库名称<span style="color: var(--error-color)">*</span></h3>
<a-input v-model:value="newDatabase.name" placeholder="新建数据库名称" /> <a-input v-model:value="newDatabase.name" placeholder="新建知识库名称" />
<h3 style="margin-top: 20px;">数据库描述</h3> <h3 style="margin-top: 20px;">知识库描述</h3>
<p style="color: var(--gray-700); font-size: 14px;">在智能体流程中这里的描述会作为工具的描述智能体会根据知识库的标题和描述来选择合适的工具所以这里描述的越详细智能体越容易选择到合适的工具</p>
<a-textarea <a-textarea
v-model:value="newDatabase.description" v-model:value="newDatabase.description"
placeholder="新建数据库描述" placeholder="新建知识库描述"
:auto-size="{ minRows: 2, maxRows: 5 }" :auto-size="{ minRows: 5, maxRows: 10 }"
/> />
<!-- <h3 style="margin-top: 20px;">向量维度</h3> <!-- <h3 style="margin-top: 20px;">向量维度</h3>
<p>必须与向量模型 {{ configStore.config.embed_model }} 一致</p> <p>必须与向量模型 {{ configStore.config.embed_model }} 一致</p>
@ -31,7 +32,7 @@
<div class="top"> <div class="top">
<div class="icon"><PlusOutlined /></div> <div class="icon"><PlusOutlined /></div>
<div class="info"> <div class="info">
<h3>新建数据</h3> <h3>新建知识</h3>
</div> </div>
</div> </div>
<p>导入您自己的文本数据或通过Webhook实时写入数据以增强 LLM 的上下文</p> <p>导入您自己的文本数据或通过Webhook实时写入数据以增强 LLM 的上下文</p>