From b9a79f73618aed1ae7f89140e3e09b72d31b2f41 Mon Sep 17 00:00:00 2001 From: martsforever Date: Thu, 11 Sep 2025 00:45:19 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BC=81=E4=B8=9A=E5=86=85=E9=83=A8?= =?UTF-8?q?=E6=96=87=E6=A1=A3=E6=A3=80=E7=B4=A2=E5=B7=A5=E5=85=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/tools/tool_list.py | 2 ++ app/tools/tool_project_analysis.py | 2 +- app/tools/tool_project_report.py | 2 +- app/tools/tool_query_direct_subordinates.py | 4 +++- app/tools/tool_query_projects.py | 4 +++- app/tools/tool_retrieve_documents.py | 18 ++++++++++++++++++ 6 files changed, 28 insertions(+), 4 deletions(-) create mode 100644 app/tools/tool_retrieve_documents.py diff --git a/app/tools/tool_list.py b/app/tools/tool_list.py index de6350e..2aa535b 100644 --- a/app/tools/tool_list.py +++ b/app/tools/tool_list.py @@ -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, ] diff --git a/app/tools/tool_project_analysis.py b/app/tools/tool_project_analysis.py index 54209a8..cf8baf0 100644 --- a/app/tools/tool_project_analysis.py +++ b/app/tools/tool_project_analysis.py @@ -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, diff --git a/app/tools/tool_project_report.py b/app/tools/tool_project_report.py index 25bdca9..71bebaf 100644 --- a/app/tools/tool_project_report.py +++ b/app/tools/tool_project_report.py @@ -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, diff --git a/app/tools/tool_query_direct_subordinates.py b/app/tools/tool_query_direct_subordinates.py index 69e3ef1..4e97ebd 100644 --- a/app/tools/tool_query_direct_subordinates.py +++ b/app/tools/tool_query_direct_subordinates.py @@ -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 [ "已经查询完毕", { diff --git a/app/tools/tool_query_projects.py b/app/tools/tool_query_projects.py index e573725..cf3c8d8 100644 --- a/app/tools/tool_query_projects.py +++ b/app/tools/tool_query_projects.py @@ -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: diff --git a/app/tools/tool_retrieve_documents.py b/app/tools/tool_retrieve_documents.py new file mode 100644 index 0000000..2cbb5ff --- /dev/null +++ b/app/tools/tool_retrieve_documents.py @@ -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