feat: 企业内部文档检索工具
This commit is contained in:
@@ -6,6 +6,7 @@ from app.tools.tool_project_analysis import tool_project_analysis
|
||||
from app.tools.tool_project_report import tool_project_report
|
||||
from app.tools.tool_query_direct_subordinates import tool_query_direct_subordinates
|
||||
from app.tools.tool_query_projects import tool_query_projects
|
||||
from app.tools.tool_retrieve_documents import tool_retrieve_documents
|
||||
|
||||
tool_list = [
|
||||
tool_book_hotel,
|
||||
@@ -16,4 +17,5 @@ tool_list = [
|
||||
tool_query_projects,
|
||||
tool_project_analysis,
|
||||
tool_project_report,
|
||||
tool_retrieve_documents,
|
||||
]
|
||||
|
||||
@@ -16,7 +16,7 @@ from app.utils.db_utils import async_session
|
||||
name_or_callable="tool_project_analysis",
|
||||
description="做项目成本分析报告"
|
||||
)
|
||||
async def tool_project_analysis(project_name: Annotated[str, "项目名称"]) -> float:
|
||||
async def tool_project_analysis(project_name: Annotated[str, "项目名称"]) -> str:
|
||||
async with async_session() as session:
|
||||
project_cls: ProjectModel = await ProjectService.query_item(
|
||||
session=session,
|
||||
|
||||
@@ -20,7 +20,7 @@ async def tool_project_report(
|
||||
project_name: Annotated[str, "项目名称"],
|
||||
start_time: Annotated[Optional[str], "开始时间(可选参数)格式为YYYY-MM-DD"] = None,
|
||||
end_time: Annotated[Optional[str], "结束时间(可选参数)格式为YYYY-MM-DD"] = None
|
||||
) -> float:
|
||||
) -> str:
|
||||
async with async_session() as session:
|
||||
project_cls: ProjectModel = await ProjectService.query_item(
|
||||
session=session,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import List, Any
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
|
||||
@@ -5,7 +7,7 @@ from langchain_core.tools import tool
|
||||
name_or_callable="tool_query_direct_subordinates",
|
||||
description="查询直接下级员工信息"
|
||||
)
|
||||
def tool_query_direct_subordinates() -> float:
|
||||
def tool_query_direct_subordinates() -> List[Any]:
|
||||
return [
|
||||
"已经查询完毕",
|
||||
{
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import List, Any
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.tools import tool
|
||||
|
||||
@@ -11,7 +13,7 @@ from app.utils.db_utils import async_session
|
||||
name_or_callable="tool_query_projects",
|
||||
description="查询参与的项目信息"
|
||||
)
|
||||
async def tool_query_projects(config: RunnableConfig) -> float:
|
||||
async def tool_query_projects(config: RunnableConfig) -> List[Any]:
|
||||
user_id = config.get('configurable').get('user_id')
|
||||
|
||||
async with async_session() as session:
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
from typing import Annotated
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
from app.utils.milvus_utils import milvus_service, KnowledgeQueryParam
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="tool_retrieve_documents",
|
||||
description="检索企业内部文档工具"
|
||||
)
|
||||
async def tool_retrieve_documents(question: Annotated[str, "用户的问题"]) -> str:
|
||||
result = await milvus_service.async_search(param=KnowledgeQueryParam(
|
||||
question=question,
|
||||
kb_code="online_document",
|
||||
top_k=10,
|
||||
))
|
||||
return result.answer
|
||||
Reference in New Issue
Block a user