feat: 查询参与项目的信息
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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]},
|
||||
}
|
||||
]
|
||||
Reference in New Issue
Block a user