From 1697324bf45713679d4e34c1529e2c26829d665c Mon Sep 17 00:00:00 2001 From: martsforever Date: Thu, 11 Sep 2025 00:32:59 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9F=A5=E8=AF=A2=E5=91=A8=E6=8A=A5?= =?UTF-8?q?=E6=97=A5=E6=8A=A5=E7=9A=84=E5=B7=A5=E5=85=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/controller/add_langgraph_chat_route.py | 10 ++- app/tools/tool_list.py | 2 + app/tools/tool_project_report.py | 74 +++++++++++++++++++--- 3 files changed, 75 insertions(+), 11 deletions(-) diff --git a/app/controller/add_langgraph_chat_route.py b/app/controller/add_langgraph_chat_route.py index ae36ba7..0d6cfc0 100644 --- a/app/controller/add_langgraph_chat_route.py +++ b/app/controller/add_langgraph_chat_route.py @@ -13,6 +13,7 @@ from starlette.responses import StreamingResponse from app.model.ConversationModel import ConversationService from app.tools.tool_list import tool_list +from app.tools.tool_project_report import tool_project_report from app.utils.db_utils import AsyncSessionDep from app.utils.llm_utils import create_llm from app.utils.postgres_checkpointer import PostgresCheckpointerManager, AsyncPostgresSaverDep @@ -43,6 +44,7 @@ class ChatAgent: prompt=""" - 你是一名擅长使用工具的智能助手,你需要根据用户问题来进行回答,请使用中文进行回答。 - 当用户问题需要调用工具时再调用工具,否则按照你的知识来回答问题。 + - 你每次只能调用一个工具 - 特别注意,特别注意,特别注意,如果工具的返回结果是一个数组,你只能回复“工具已经执行完毕” """ ) @@ -63,8 +65,12 @@ class ChatAgent: def add_langgraph_chat_route(app: FastAPI): @app.post("/project/analysis") async def project_analysis(body: dict): - project_name = body.get('name') - return {} + + return await tool_project_report.ainvoke({ + "project_name": body.get('project_name'), + "start_time": body.get("start_time", None), + "end_time": body.get("end_time", None), + }) # 流式对话接口 @app.post("/langgraph/stream") diff --git a/app/tools/tool_list.py b/app/tools/tool_list.py index a766540..de6350e 100644 --- a/app/tools/tool_list.py +++ b/app/tools/tool_list.py @@ -3,6 +3,7 @@ from app.tools.tool_book_hotel import tool_book_hotel from app.tools.tool_get_datetime import tool_get_datetime from app.tools.tool_multiple import tool_multiply 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 @@ -14,4 +15,5 @@ tool_list = [ tool_query_direct_subordinates, tool_query_projects, tool_project_analysis, + tool_project_report, ] diff --git a/app/tools/tool_project_report.py b/app/tools/tool_project_report.py index 08cedfb..25bdca9 100644 --- a/app/tools/tool_project_report.py +++ b/app/tools/tool_project_report.py @@ -1,15 +1,71 @@ +import json +from typing import Annotated, List, Optional + from langchain_core.tools import tool +from sqlalchemy.orm import selectinload +from sqlmodel import select + +from app.model.KnowledgeDoc import KnowledgeDocModel +from app.model.ProjectModel import ProjectModel, ProjectService +from app.model.RelProjUserModel import RelProjUserService, RelProjUserModel +from app.utils.PageQueryParams import PageQueryParams +from app.utils.db_utils import async_session, AsyncSessionDep @tool( name_or_callable="tool_project_report", - description="tool_project_report" + description="查询项目日报周报的工具" ) -def tool_project_analysis() -> float: - return [ - "", - { - "component": "DirectSubordinates", - "props": {}, - } - ] +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: + async with async_session() as session: + project_cls: ProjectModel = await ProjectService.query_item( + session=session, + row_dict={"name": project_name} + ) + if not project_cls: + return f"找不到名为「{project_name}」的项目" + + # 查询所有所属的项目成员 + query_result = await RelProjUserService.query_list(session=session, query_param=PageQueryParams(all=True, filters={"proj_id": project_cls.id})) + rel_list: List[RelProjUserModel] = query_result[0] + member_id_list: List[str] = [item.user_id for item in rel_list] + + # member_list = await UserService.query_list(session=session, query_param=PageQueryParams(all=True, filters={"id": member_id_list})) + + doc_list = await query_users_reports(session=session, user_id_list=member_id_list, start_time=start_time, end_time=end_time) + + user_reports = {} + + for doc in doc_list: + user_full_name = doc.creator.full_name + if user_full_name not in user_reports: + user_reports[user_full_name] = [] + user_reports[user_full_name].append({ + "datetime": doc.created_at.strftime("%Y-%m-%d") if doc.created_at else None, + "content": doc.content, + }) + + return json.dumps(user_reports, ensure_ascii=False) + + +async def query_users_reports(session: AsyncSessionDep, user_id_list: List[str], start_time: Optional[str], end_time: Optional[str]): + filter_parent_codes = [f"report_{user_id}" for user_id in user_id_list] + + query = select(KnowledgeDocModel).options(selectinload(KnowledgeDocModel.creator_relationship)) + + query = query.where(KnowledgeDocModel.parent_code.in_(filter_parent_codes)) + if start_time: + query = query.where(KnowledgeDocModel.created_at >= (start_time + " 00:00:00")) + if end_time: + query = query.where(KnowledgeDocModel.created_at <= (end_time + " 23:59:59")) + + query = query.order_by(KnowledgeDocModel.created_at.desc()) + + result = await session.execute(query) + query_cls_list: List[KnowledgeDocModel] = result.scalars().all() + + return query_cls_list