feat: 完成实现通过对话实现酒店预定的功能

This commit is contained in:
martsforever
2025-09-04 18:55:53 +08:00
parent 88963dcca0
commit 9fe90a69f7
+42 -10
View File
@@ -3,7 +3,8 @@ import json
import time import time
from typing import Union, Annotated, TypedDict from typing import Union, Annotated, TypedDict
from fastapi import FastAPI import httpx
from fastapi import FastAPI, Depends
from langchain_core.messages import AIMessage, ToolMessage from langchain_core.messages import AIMessage, ToolMessage
from langchain_core.runnables import RunnableConfig from langchain_core.runnables import RunnableConfig
from langchain_core.tools import tool from langchain_core.tools import tool
@@ -13,6 +14,8 @@ from langgraph.types import interrupt, Command
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from starlette.responses import StreamingResponse from starlette.responses import StreamingResponse
from app.config.env import env
from app.controller.add_user_route import oauth2_scheme
from app.model.ConversationModel import ConversationService from app.model.ConversationModel import ConversationService
from app.utils.db_utils import AsyncSessionDep from app.utils.db_utils import AsyncSessionDep
from app.utils.llm_utils import create_llm from app.utils.llm_utils import create_llm
@@ -66,21 +69,42 @@ class BookHotelData(TypedDict):
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具") @tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
def tool_book_hotel(config: RunnableConfig) -> str: async def tool_book_hotel(config: RunnableConfig) -> str:
"""
触发中断,输入酒店预定表单信息
"""
book_hotel_data: Union[BookHotelData, str] = interrupt({"title": "请输入酒店预定信息", "formCode": "bookHotel"}) book_hotel_data: Union[BookHotelData, str] = interrupt({"title": "请输入酒店预定信息", "formCode": "bookHotel"})
if book_hotel_data == "N": if book_hotel_data == "N":
return f"用户选择取消预定酒店" return f"用户选择取消预定酒店"
hotel_id = book_hotel_data.get('hotel_id') hotel_id = book_hotel_data.get('hotel_id')
proj_id = book_hotel_data.get('proj_id') proj_id = book_hotel_data.get('proj_id')
user_id = book_hotel_data.get('user_id') user_id = book_hotel_data.get('user_id')
token = config['configurable']['token']
print("hotel_id", hotel_id) # print("================ tool_book_hotel ================")
print("proj_id", proj_id) # print("hotel_id", hotel_id)
print("user_id", user_id) # print("proj_id", proj_id)
print("config", config) # print("user_id", user_id)
return "预定成功" async with httpx.AsyncClient() as client:
response = await client.post(
url=f"http://localhost:{env.server_port}/book_hotel",
headers={
"Authorization": f"Bearer {token}"
},
json={
"hotel_id": hotel_id,
"user_id": user_id,
"proj_id": proj_id,
},
)
response_data = response.json()
# print("================ response_data ================")
# print(response_data)
return "预定成功," + response_data.get('message')
class ChatMessage(BaseModel): class ChatMessage(BaseModel):
@@ -133,12 +157,16 @@ class ChatAgent:
def add_langgraph_chat_route(app: FastAPI): def add_langgraph_chat_route(app: FastAPI):
# 流式对话接口 # 流式对话接口
@app.post("/langgraph/stream") @app.post("/langgraph/stream")
async def langgraph_stream(body: dict, checkpointer: AsyncPostgresSaverDep): async def langgraph_stream(
body: dict,
token: str = Depends(oauth2_scheme),
):
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::") print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
print(body) print(body)
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
print("stream_input", stream_input) print("stream_input", stream_input)
print("stream_config", stream_config) print("stream_config", stream_config)
@@ -207,13 +235,17 @@ def add_langgraph_chat_route(app: FastAPI):
# 恢复中断接口 # 恢复中断接口
@app.post("/langgraph/chat_resume/{thread_id}") @app.post("/langgraph/chat_resume/{thread_id}")
async def langgraph_chat(body: dict, thread_id: str): async def langgraph_chat(
body: dict,
thread_id: str,
token: str = Depends(oauth2_scheme),
):
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}} config={"configurable": {"thread_id": thread_id, "token": token}}
) )
return { return {
**graph_state, **graph_state,