refactor: 重构知识库工具模块,优化工具获取逻辑并移除冗余代码
This commit is contained in:
parent
2020438152
commit
09577b9b65
@ -1,7 +1,11 @@
|
|||||||
from deepagents.middleware.filesystem import FilesystemMiddleware
|
from deepagents.middleware.filesystem import FilesystemMiddleware
|
||||||
from langchain.agents import create_agent
|
from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
||||||
from langchain.agents.middleware import ModelRetryMiddleware
|
|
||||||
|
|
||||||
|
from langchain.agents import create_agent
|
||||||
|
from langchain.agents.middleware import (
|
||||||
|
TodoListMiddleware,
|
||||||
|
ModelRetryMiddleware,
|
||||||
|
)
|
||||||
from src.agents.common import BaseAgent, load_chat_model
|
from src.agents.common import BaseAgent, load_chat_model
|
||||||
from src.agents.common.backends import create_agent_composite_backend
|
from src.agents.common.backends import create_agent_composite_backend
|
||||||
from src.agents.common.middlewares import (
|
from src.agents.common.middlewares import (
|
||||||
@ -19,7 +23,7 @@ def _create_fs_backend(rt):
|
|||||||
class ChatbotAgent(BaseAgent):
|
class ChatbotAgent(BaseAgent):
|
||||||
name = "智能体助手"
|
name = "智能体助手"
|
||||||
description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。"
|
description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。"
|
||||||
capabilities = ["file_upload", "files"] # 支持文件上传功能
|
capabilities = ["file_upload", "files", "todo"] # 支持文件上传功能
|
||||||
|
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@ -34,13 +38,15 @@ class ChatbotAgent(BaseAgent):
|
|||||||
# 使用 create_agent 创建智能体
|
# 使用 create_agent 创建智能体
|
||||||
# 注意:tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置
|
# 注意:tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置
|
||||||
graph = create_agent(
|
graph = create_agent(
|
||||||
model=load_chat_model(context.model),
|
model=load_chat_model(fully_specified_name=context.model),
|
||||||
system_prompt=context.system_prompt,
|
system_prompt=context.system_prompt,
|
||||||
middleware=[
|
middleware=[
|
||||||
save_attachments_to_fs, # 附件注入提示词
|
save_attachments_to_fs, # 附件注入提示词
|
||||||
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
|
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
|
||||||
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/知识库/MCP/提示词)
|
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/知识库/MCP/提示词)
|
||||||
ModelRetryMiddleware(), # 模型重试中间件
|
ModelRetryMiddleware(), # 模型重试中间件
|
||||||
|
TodoListMiddleware(),
|
||||||
|
PatchToolCallsMiddleware(),
|
||||||
],
|
],
|
||||||
checkpointer=await self._get_checkpointer(),
|
checkpointer=await self._get_checkpointer(),
|
||||||
)
|
)
|
||||||
|
|||||||
3
src/agents/common/toolkits/kbs/__init__.py
Normal file
3
src/agents/common/toolkits/kbs/__init__.py
Normal file
@ -0,0 +1,3 @@
|
|||||||
|
from .tools import get_kb_based_tools
|
||||||
|
|
||||||
|
__all__ = ["get_kb_based_tools"]
|
||||||
159
src/agents/common/toolkits/kbs/tools.py
Normal file
159
src/agents/common/toolkits/kbs/tools.py
Normal file
@ -0,0 +1,159 @@
|
|||||||
|
"""知识库工具模块"""
|
||||||
|
import inspect
|
||||||
|
import traceback
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain_core.tools import StructuredTool
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from src import knowledge_base
|
||||||
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
|
class KnowledgeRetrieverModel(BaseModel):
|
||||||
|
query_text: str | None = Field(
|
||||||
|
default=None,
|
||||||
|
description=(
|
||||||
|
"查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
operation: str = Field(
|
||||||
|
default="search",
|
||||||
|
description=(
|
||||||
|
"操作类型:'search' 表示检索知识库内容,'get_mindmap' 表示获取知识库的思维导图结构。"
|
||||||
|
"当用户询问知识库的整体结构、文件分类、知识架构时,使用 'get_mindmap'。"
|
||||||
|
"当用户需要查询具体内容时,使用 'search'。"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CommonKnowledgeRetriever(KnowledgeRetrieverModel):
|
||||||
|
"""Common knowledge retriever model."""
|
||||||
|
|
||||||
|
file_name: str | None = Field(
|
||||||
|
default=None,
|
||||||
|
description=(
|
||||||
|
"(非必要不启用此参数,留空即可)当操作类型为 'search' 且已经读取思维导图之后,可以指定文件关键词,支持模糊匹配。\n"
|
||||||
|
"仅当检索结果过多且不相关,需要进一步缩小范围时使用。"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_kb_based_tools(db_names: list[str] | None = None) -> list:
|
||||||
|
"""获取所有知识库基于的工具"""
|
||||||
|
# 获取所有知识库
|
||||||
|
kb_tools = []
|
||||||
|
retrievers = knowledge_base.get_retrievers()
|
||||||
|
if db_names is None:
|
||||||
|
db_ids = None
|
||||||
|
else:
|
||||||
|
db_ids = [kb_id for kb_id, kb in retrievers.items() if kb["name"] in db_names]
|
||||||
|
|
||||||
|
def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]):
|
||||||
|
"""创建检索器包装函数的工厂函数,避免闭包变量捕获问题"""
|
||||||
|
|
||||||
|
async def async_retriever_wrapper(
|
||||||
|
query_text: str, operation: str = "search", file_name: str | None = None
|
||||||
|
) -> Any:
|
||||||
|
"""异步检索器包装函数,支持检索和获取思维导图"""
|
||||||
|
|
||||||
|
# 获取思维导图
|
||||||
|
if operation == "get_mindmap":
|
||||||
|
try:
|
||||||
|
logger.debug(f"Getting mindmap for database {db_id}")
|
||||||
|
|
||||||
|
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||||
|
|
||||||
|
kb_repo = KnowledgeBaseRepository()
|
||||||
|
kb = await kb_repo.get_by_id(db_id)
|
||||||
|
|
||||||
|
if kb is None:
|
||||||
|
return f"知识库 {retriever_info['name']} 不存在"
|
||||||
|
|
||||||
|
mindmap_data = kb.mindmap
|
||||||
|
|
||||||
|
if not mindmap_data:
|
||||||
|
return f"知识库 {retriever_info['name']} 还没有生成思维导图。"
|
||||||
|
|
||||||
|
# 将思维导图数据转换为文本格式,便于AI理解
|
||||||
|
def mindmap_to_text(node, level=0):
|
||||||
|
"""递归将思维导图JSON转换为层级文本"""
|
||||||
|
indent = " " * level
|
||||||
|
text = f"{indent}- {node.get('content', '')}\n"
|
||||||
|
for child in node.get("children", []):
|
||||||
|
text += mindmap_to_text(child, level + 1)
|
||||||
|
return text
|
||||||
|
|
||||||
|
mindmap_text = f"知识库 {retriever_info['name']} 的思维导图结构:\n\n"
|
||||||
|
mindmap_text += mindmap_to_text(mindmap_data)
|
||||||
|
|
||||||
|
logger.debug(f"Successfully retrieved mindmap for {db_id}")
|
||||||
|
return mindmap_text
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting mindmap for {db_id}: {e}")
|
||||||
|
return f"获取思维导图失败: {str(e)}"
|
||||||
|
|
||||||
|
# 默认:检索知识库
|
||||||
|
retriever = retriever_info["retriever"]
|
||||||
|
try:
|
||||||
|
logger.debug(f"Retrieving from database {db_id} with query: {query_text}")
|
||||||
|
kwargs = {}
|
||||||
|
if file_name:
|
||||||
|
kwargs["file_name"] = file_name
|
||||||
|
|
||||||
|
if inspect.iscoroutinefunction(retriever):
|
||||||
|
result = await retriever(query_text, **kwargs)
|
||||||
|
else:
|
||||||
|
result = retriever(query_text, **kwargs)
|
||||||
|
logger.debug(f"Retrieved {len(result) if isinstance(result, list) else 'N/A'} results from {db_id}")
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error in retriever {db_id}: {e}")
|
||||||
|
return f"检索失败: {str(e)}"
|
||||||
|
|
||||||
|
return async_retriever_wrapper
|
||||||
|
|
||||||
|
for db_id, retrieve_info in retrievers.items():
|
||||||
|
if db_ids is not None and db_id not in db_ids:
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 构建工具描述
|
||||||
|
description = (
|
||||||
|
f"使用 {retrieve_info['name']} 知识库的多功能工具。\n"
|
||||||
|
f"知识库描述:{retrieve_info['description'] or '没有描述。'}\n\n"
|
||||||
|
f"支持的操作:\n"
|
||||||
|
f"1. 'search' - 检索知识库内容:根据关键词查询相关文档片段\n"
|
||||||
|
f"2. 'get_mindmap' - 获取思维导图:查看知识库的整体结构和文件分类\n\n"
|
||||||
|
f"使用建议:\n"
|
||||||
|
f"- 需要查询具体内容时,使用 operation='search'\n"
|
||||||
|
f"- 想了解知识库结构、文件分类时,使用 operation='get_mindmap'"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 使用工厂函数创建检索器包装函数,避免闭包问题
|
||||||
|
retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info)
|
||||||
|
|
||||||
|
safename = retrieve_info["name"].replace(" ", "_")[:20]
|
||||||
|
|
||||||
|
args_schema = KnowledgeRetrieverModel
|
||||||
|
if retrieve_info["metadata"]["kb_type"] in ["milvus"]:
|
||||||
|
args_schema = CommonKnowledgeRetriever
|
||||||
|
|
||||||
|
# 使用 StructuredTool.from_function 创建异步工具
|
||||||
|
tool = StructuredTool.from_function(
|
||||||
|
coroutine=retriever_wrapper,
|
||||||
|
name=safename,
|
||||||
|
description=description,
|
||||||
|
args_schema=args_schema,
|
||||||
|
metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]},
|
||||||
|
)
|
||||||
|
|
||||||
|
kb_tools.append(tool)
|
||||||
|
# logger.debug(f"Successfully created tool {tool_id} for database {db_id}")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to create tool for database {db_id}: {e}, \n{traceback.format_exc()}")
|
||||||
|
continue
|
||||||
|
|
||||||
|
return kb_tools
|
||||||
@ -1,4 +1,3 @@
|
|||||||
import asyncio
|
|
||||||
import os
|
import os
|
||||||
import traceback
|
import traceback
|
||||||
import uuid
|
import uuid
|
||||||
@ -6,11 +5,10 @@ from typing import Annotated, Any
|
|||||||
|
|
||||||
import requests
|
import requests
|
||||||
from langchain.tools import tool
|
from langchain.tools import tool
|
||||||
from langchain_core.tools import StructuredTool
|
|
||||||
from langgraph.types import interrupt
|
from langgraph.types import interrupt
|
||||||
from pydantic import BaseModel, Field
|
|
||||||
|
|
||||||
from src import config, graph_base, knowledge_base
|
from src import config, graph_base
|
||||||
|
from src.agents.common.toolkits.kbs import get_kb_based_tools
|
||||||
from src.services.mcp_service import get_enabled_mcp_tools
|
from src.services.mcp_service import get_enabled_mcp_tools
|
||||||
from src.storage.minio import aupload_file_to_minio
|
from src.storage.minio import aupload_file_to_minio
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
@ -147,148 +145,6 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge
|
|||||||
return f"知识图谱查询失败: {str(e)}"
|
return f"知识图谱查询失败: {str(e)}"
|
||||||
|
|
||||||
|
|
||||||
class KnowledgeRetrieverModel(BaseModel):
|
|
||||||
query_text: str = Field(
|
|
||||||
description=(
|
|
||||||
"查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
operation: str = Field(
|
|
||||||
default="search",
|
|
||||||
description=(
|
|
||||||
"操作类型:'search' 表示检索知识库内容,'get_mindmap' 表示获取知识库的思维导图结构。"
|
|
||||||
"当用户询问知识库的整体结构、文件分类、知识架构时,使用 'get_mindmap'。"
|
|
||||||
"当用户需要查询具体内容时,使用 'search'。"
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class CommonKnowledgeRetriever(KnowledgeRetrieverModel):
|
|
||||||
"""Common knowledge retriever model."""
|
|
||||||
|
|
||||||
file_name: str = Field(description="限定文件名称,当操作类型为 'search' 时,可以指定文件名称,支持模糊匹配")
|
|
||||||
|
|
||||||
|
|
||||||
def get_kb_based_tools(db_names: list[str] | None = None) -> list:
|
|
||||||
"""获取所有知识库基于的工具"""
|
|
||||||
# 获取所有知识库
|
|
||||||
kb_tools = []
|
|
||||||
retrievers = knowledge_base.get_retrievers()
|
|
||||||
if db_names is None:
|
|
||||||
db_ids = None
|
|
||||||
else:
|
|
||||||
db_ids = [kb_id for kb_id, kb in retrievers.items() if kb["name"] in db_names]
|
|
||||||
|
|
||||||
def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]):
|
|
||||||
"""创建检索器包装函数的工厂函数,避免闭包变量捕获问题"""
|
|
||||||
|
|
||||||
async def async_retriever_wrapper(
|
|
||||||
query_text: str, operation: str = "search", file_name: str | None = None
|
|
||||||
) -> Any:
|
|
||||||
"""异步检索器包装函数,支持检索和获取思维导图"""
|
|
||||||
|
|
||||||
# 获取思维导图
|
|
||||||
if operation == "get_mindmap":
|
|
||||||
try:
|
|
||||||
logger.debug(f"Getting mindmap for database {db_id}")
|
|
||||||
|
|
||||||
from src.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
|
||||||
|
|
||||||
kb_repo = KnowledgeBaseRepository()
|
|
||||||
kb = await kb_repo.get_by_id(db_id)
|
|
||||||
|
|
||||||
if kb is None:
|
|
||||||
return f"知识库 {retriever_info['name']} 不存在"
|
|
||||||
|
|
||||||
mindmap_data = kb.mindmap
|
|
||||||
|
|
||||||
if not mindmap_data:
|
|
||||||
return f"知识库 {retriever_info['name']} 还没有生成思维导图。"
|
|
||||||
|
|
||||||
# 将思维导图数据转换为文本格式,便于AI理解
|
|
||||||
def mindmap_to_text(node, level=0):
|
|
||||||
"""递归将思维导图JSON转换为层级文本"""
|
|
||||||
indent = " " * level
|
|
||||||
text = f"{indent}- {node.get('content', '')}\n"
|
|
||||||
for child in node.get("children", []):
|
|
||||||
text += mindmap_to_text(child, level + 1)
|
|
||||||
return text
|
|
||||||
|
|
||||||
mindmap_text = f"知识库 {retriever_info['name']} 的思维导图结构:\n\n"
|
|
||||||
mindmap_text += mindmap_to_text(mindmap_data)
|
|
||||||
|
|
||||||
logger.debug(f"Successfully retrieved mindmap for {db_id}")
|
|
||||||
return mindmap_text
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error getting mindmap for {db_id}: {e}")
|
|
||||||
return f"获取思维导图失败: {str(e)}"
|
|
||||||
|
|
||||||
# 默认:检索知识库
|
|
||||||
retriever = retriever_info["retriever"]
|
|
||||||
try:
|
|
||||||
logger.debug(f"Retrieving from database {db_id} with query: {query_text}")
|
|
||||||
kwargs = {}
|
|
||||||
if file_name:
|
|
||||||
kwargs["file_name"] = file_name
|
|
||||||
|
|
||||||
if asyncio.iscoroutinefunction(retriever):
|
|
||||||
result = await retriever(query_text, **kwargs)
|
|
||||||
else:
|
|
||||||
result = retriever(query_text, **kwargs)
|
|
||||||
logger.debug(f"Retrieved {len(result) if isinstance(result, list) else 'N/A'} results from {db_id}")
|
|
||||||
return result
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error in retriever {db_id}: {e}")
|
|
||||||
return f"检索失败: {str(e)}"
|
|
||||||
|
|
||||||
return async_retriever_wrapper
|
|
||||||
|
|
||||||
for db_id, retrieve_info in retrievers.items():
|
|
||||||
if db_ids is not None and db_id not in db_ids:
|
|
||||||
continue
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 构建工具描述
|
|
||||||
description = (
|
|
||||||
f"使用 {retrieve_info['name']} 知识库的多功能工具。\n"
|
|
||||||
f"知识库描述:{retrieve_info['description'] or '没有描述。'}\n\n"
|
|
||||||
f"支持的操作:\n"
|
|
||||||
f"1. 'search' - 检索知识库内容:根据关键词查询相关文档片段\n"
|
|
||||||
f"2. 'get_mindmap' - 获取思维导图:查看知识库的整体结构和文件分类\n\n"
|
|
||||||
f"使用建议:\n"
|
|
||||||
f"- 需要查询具体内容时,使用 operation='search'\n"
|
|
||||||
f"- 想了解知识库结构、文件分类时,使用 operation='get_mindmap'"
|
|
||||||
)
|
|
||||||
|
|
||||||
# 使用工厂函数创建检索器包装函数,避免闭包问题
|
|
||||||
retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info)
|
|
||||||
|
|
||||||
safename = retrieve_info["name"].replace(" ", "_")[:20]
|
|
||||||
|
|
||||||
args_schema = KnowledgeRetrieverModel
|
|
||||||
if retrieve_info["metadata"]["kb_type"] in ["milvus"]:
|
|
||||||
args_schema = CommonKnowledgeRetriever
|
|
||||||
|
|
||||||
# 使用 StructuredTool.from_function 创建异步工具
|
|
||||||
tool = StructuredTool.from_function(
|
|
||||||
coroutine=retriever_wrapper,
|
|
||||||
name=safename,
|
|
||||||
description=description,
|
|
||||||
args_schema=args_schema,
|
|
||||||
metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]},
|
|
||||||
)
|
|
||||||
|
|
||||||
kb_tools.append(tool)
|
|
||||||
# logger.debug(f"Successfully created tool {tool_id} for database {db_id}")
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Failed to create tool for database {db_id}: {e}, \n{traceback.format_exc()}")
|
|
||||||
continue
|
|
||||||
|
|
||||||
return kb_tools
|
|
||||||
|
|
||||||
|
|
||||||
def gen_tool_info(tools) -> list[dict[str, Any]]:
|
def gen_tool_info(tools) -> list[dict[str, Any]]:
|
||||||
"""获取所有工具的信息(用于前端展示)"""
|
"""获取所有工具的信息(用于前端展示)"""
|
||||||
tools_info = []
|
tools_info = []
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user