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 import time
from typing import Union from typing import Union
from fastapi import FastAPI, Depends from fastapi import FastAPI
from langchain_core.messages import AIMessage, ToolMessage from langchain_core.messages import AIMessage, ToolMessage
from langgraph.graph.state import CompiledStateGraph from langgraph.graph.state import CompiledStateGraph
from langgraph.prebuilt import create_react_agent from langgraph.prebuilt import create_react_agent
from langgraph.types import Command from langgraph.types import Command
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from starlette.requests import Request
from starlette.responses import StreamingResponse from starlette.responses import StreamingResponse
from app.controller.add_user_route import oauth2_scheme
from app.model.ConversationModel import ConversationService from app.model.ConversationModel import ConversationService
from app.tools.tool_list import tool_list from app.tools.tool_list import tool_list
from app.utils.db_utils import AsyncSessionDep from app.utils.db_utils import AsyncSessionDep
@@ -65,14 +65,17 @@ def add_langgraph_chat_route(app: FastAPI):
@app.post("/langgraph/stream") @app.post("/langgraph/stream")
async def langgraph_stream( async def langgraph_stream(
body: dict, body: dict,
token: str = Depends(oauth2_scheme), request: Request,
): ):
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::") print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
print(body) print(body)
print(request.state.user)
print(request.state.token)
stream_input = body.get('input') stream_input = body.get('input')
stream_config = body.get('config') 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_input", stream_input)
print("stream_config", stream_config) print("stream_config", stream_config)
@@ -104,6 +107,8 @@ def add_langgraph_chat_route(app: FastAPI):
if emit_chunk['stream_type'] == "messages": if emit_chunk['stream_type'] == "messages":
# messages模式流式输出,此时 chunk[1][0] 为AIMessageChunk # messages模式流式输出,此时 chunk[1][0] 为AIMessageChunk
chunk_message = chunk[1][0] chunk_message = chunk[1][0]
if not chunk_message.content:
continue
else: else:
# updates模式流式输出 # updates模式流式输出
for k, v in chunk[1].items(): for k, v in chunk[1].items():
@@ -144,14 +149,18 @@ def add_langgraph_chat_route(app: FastAPI):
async def langgraph_chat( async def langgraph_chat(
body: dict, body: dict,
thread_id: str, thread_id: str,
token: str = Depends(oauth2_scheme), request: Request,
): ):
graph = await ChatAgent.get_agent() graph = await ChatAgent.get_agent()
chat_state = await ChatAgent.get_chat_state(thread_id) chat_state = await ChatAgent.get_chat_state(thread_id)
chat_history_list = chat_state.get('messages') chat_history_list = chat_state.get('messages')
graph_state = await graph.ainvoke( graph_state = await graph.ainvoke(
Command(resume=body.get('resume_data')), 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 { return {
**graph_state, **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_get_datetime import tool_get_datetime
from app.tools.tool_multiple import tool_multiply 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_direct_subordinates import tool_query_direct_subordinates
from app.tools.tool_query_projects import tool_query_projects
tool_list = [ tool_list = [
tool_book_hotel, tool_book_hotel,
@@ -10,4 +11,5 @@ tool_list = [
tool_add, tool_add,
tool_multiply, tool_multiply,
tool_query_direct_subordinates, 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]},
}
]