refactor: 重构知识库工具模块,优化工具获取逻辑并移除冗余代码
This commit is contained in:
parent
2020438152
commit
09577b9b65
@ -1,7 +1,11 @@
|
||||
from deepagents.middleware.filesystem import FilesystemMiddleware
|
||||
from langchain.agents import create_agent
|
||||
from langchain.agents.middleware import ModelRetryMiddleware
|
||||
from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
||||
|
||||
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.backends import create_agent_composite_backend
|
||||
from src.agents.common.middlewares import (
|
||||
@ -19,7 +23,7 @@ def _create_fs_backend(rt):
|
||||
class ChatbotAgent(BaseAgent):
|
||||
name = "智能体助手"
|
||||
description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。"
|
||||
capabilities = ["file_upload", "files"] # 支持文件上传功能
|
||||
capabilities = ["file_upload", "files", "todo"] # 支持文件上传功能
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
@ -34,13 +38,15 @@ class ChatbotAgent(BaseAgent):
|
||||
# 使用 create_agent 创建智能体
|
||||
# 注意:tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置
|
||||
graph = create_agent(
|
||||
model=load_chat_model(context.model),
|
||||
model=load_chat_model(fully_specified_name=context.model),
|
||||
system_prompt=context.system_prompt,
|
||||
middleware=[
|
||||
save_attachments_to_fs, # 附件注入提示词
|
||||
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
|
||||
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/知识库/MCP/提示词)
|
||||
ModelRetryMiddleware(), # 模型重试中间件
|
||||
TodoListMiddleware(),
|
||||
PatchToolCallsMiddleware(),
|
||||
],
|
||||
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 traceback
|
||||
import uuid
|
||||
@ -6,11 +5,10 @@ from typing import Annotated, Any
|
||||
|
||||
import requests
|
||||
from langchain.tools import tool
|
||||
from langchain_core.tools import StructuredTool
|
||||
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.storage.minio import aupload_file_to_minio
|
||||
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)}"
|
||||
|
||||
|
||||
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]]:
|
||||
"""获取所有工具的信息(用于前端展示)"""
|
||||
tools_info = []
|
||||
|
||||
Loading…
Reference in New Issue
Block a user