feat: 查询周报日报的工具
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user