From 091ccac86777e60f386a7f7e3badc62e698f9327 Mon Sep 17 00:00:00 2001 From: martsforever Date: Thu, 4 Sep 2025 22:52:24 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9F=A5=E8=AF=A2=E5=8F=82=E4=B8=8E?= =?UTF-8?q?=E9=A1=B9=E7=9B=AE=E7=9A=84=E4=BF=A1=E6=81=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/controller/add_langgraph_chat_route.py | 21 +++++++++++----- app/tools/tool_list.py | 2 ++ app/tools/tool_query_projects.py | 28 ++++++++++++++++++++++ 3 files changed, 45 insertions(+), 6 deletions(-) create mode 100644 app/tools/tool_query_projects.py diff --git a/app/controller/add_langgraph_chat_route.py b/app/controller/add_langgraph_chat_route.py index 92fabb0..14c74fd 100644 --- a/app/controller/add_langgraph_chat_route.py +++ b/app/controller/add_langgraph_chat_route.py @@ -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, diff --git a/app/tools/tool_list.py b/app/tools/tool_list.py index ca22fa6..724047a 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_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, ] diff --git a/app/tools/tool_query_projects.py b/app/tools/tool_query_projects.py new file mode 100644 index 0000000..c422839 --- /dev/null +++ b/app/tools/tool_query_projects.py @@ -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]}, + } + ]