feat: 完成实现通过对话实现酒店预定的功能
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user