feat: 查询参与项目的信息
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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