feat: 查询参与项目的信息

This commit is contained in:
martsforever
2025-09-04 22:52:24 +08:00
parent 1ada5aaffd
commit 091ccac867
3 changed files with 45 additions and 6 deletions
+15 -6
View File
@@ -2,15 +2,15 @@ import json
import time
from typing import Union
from fastapi import FastAPI, Depends
from fastapi import FastAPI
from langchain_core.messages import AIMessage, ToolMessage
from langgraph.graph.state import CompiledStateGraph
from langgraph.prebuilt import create_react_agent
from langgraph.types import Command
from pydantic import BaseModel, Field
from starlette.requests import Request
from starlette.responses import StreamingResponse
from app.controller.add_user_route import oauth2_scheme
from app.model.ConversationModel import ConversationService
from app.tools.tool_list import tool_list
from app.utils.db_utils import AsyncSessionDep
@@ -65,14 +65,17 @@ def add_langgraph_chat_route(app: FastAPI):
@app.post("/langgraph/stream")
async def langgraph_stream(
body: dict,
token: str = Depends(oauth2_scheme),
request: Request,
):
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
print(body)
print(request.state.user)
print(request.state.token)
stream_input = body.get('input')
stream_config = body.get('config')
stream_config['configurable']['token'] = token
stream_config['configurable']['token'] = request.state.token
stream_config['configurable']['user_id'] = request.state.user.id
print("stream_input", stream_input)
print("stream_config", stream_config)
@@ -104,6 +107,8 @@ def add_langgraph_chat_route(app: FastAPI):
if emit_chunk['stream_type'] == "messages":
# messages模式流式输出,此时 chunk[1][0] 为AIMessageChunk
chunk_message = chunk[1][0]
if not chunk_message.content:
continue
else:
# updates模式流式输出
for k, v in chunk[1].items():
@@ -144,14 +149,18 @@ def add_langgraph_chat_route(app: FastAPI):
async def langgraph_chat(
body: dict,
thread_id: str,
token: str = Depends(oauth2_scheme),
request: Request,
):
graph = await ChatAgent.get_agent()
chat_state = await ChatAgent.get_chat_state(thread_id)
chat_history_list = chat_state.get('messages')
graph_state = await graph.ainvoke(
Command(resume=body.get('resume_data')),
config={"configurable": {"thread_id": thread_id, "token": token}}
config={"configurable": {
"thread_id": thread_id,
"token": request.state.token,
"user_id": request.state.user.id,
}}
)
return {
**graph_state,
+2
View File
@@ -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_query_direct_subordinates import tool_query_direct_subordinates
from app.tools.tool_query_projects import tool_query_projects
tool_list = [
tool_book_hotel,
@@ -10,4 +11,5 @@ tool_list = [
tool_add,
tool_multiply,
tool_query_direct_subordinates,
tool_query_projects,
]
+28
View File
@@ -0,0 +1,28 @@
from langchain_core.runnables import RunnableConfig
from langchain_core.tools import tool
from app.model.ProjectService import ProjectService
from app.model.RelProjUserModel import RelProjUserService
from app.utils.PageQueryParams import PageQueryParams
from app.utils.db_utils import async_session
@tool(
name_or_callable="tool_query_projects",
description="查询参与的项目信息"
)
async def tool_query_projects(config: RunnableConfig) -> float:
user_id = config.get('configurable').get('user_id')
async with async_session() as session:
query_cls_list, has_next, total = await RelProjUserService.query_list(session=session, query_param=PageQueryParams(all=True, filters={"user_id": user_id}))
proj_id_list = [item.proj_id for item in query_cls_list]
proj_list, has_next, total = await ProjectService.query_list(session=session, query_param=PageQueryParams(all=True, filters={"id": proj_id_list}))
return [
"已经查询完毕",
{
"component": "ProjectList",
"props": {"projIdList": [item.id for item in proj_list]},
}
]