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
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.runnables import RunnableConfig
from langchain_core.tools import tool
@@ -13,6 +14,8 @@ from langgraph.types import interrupt, Command
from pydantic import BaseModel, Field
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.utils.db_utils import AsyncSessionDep
from app.utils.llm_utils import create_llm
@@ -66,21 +69,42 @@ class BookHotelData(TypedDict):
@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"})
if book_hotel_data == "N":
return f"用户选择取消预定酒店"
hotel_id = book_hotel_data.get('hotel_id')
proj_id = book_hotel_data.get('proj_id')
user_id = book_hotel_data.get('user_id')
token = config['configurable']['token']
print("hotel_id", hotel_id)
print("proj_id", proj_id)
print("user_id", user_id)
print("config", config)
# print("================ tool_book_hotel ================")
# print("hotel_id", hotel_id)
# print("proj_id", proj_id)
# 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):
@@ -133,12 +157,16 @@ class ChatAgent:
def add_langgraph_chat_route(app: FastAPI):
# 流式对话接口
@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(body)
stream_input = body.get('input')
stream_config = body.get('config')
stream_config['configurable']['token'] = token
print("stream_input", stream_input)
print("stream_config", stream_config)
@@ -207,13 +235,17 @@ def add_langgraph_chat_route(app: FastAPI):
# 恢复中断接口
@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()
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}}
config={"configurable": {"thread_id": thread_id, "token": token}}
)
return {
**graph_state,