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