250 lines
8.2 KiB
Python
250 lines
8.2 KiB
Python
import datetime
|
|
import json
|
|
import time
|
|
from typing import Union, Annotated
|
|
|
|
from fastapi import FastAPI
|
|
from langchain_core.messages import HumanMessage, AIMessage, ToolMessage
|
|
from langchain_core.tools import tool
|
|
from langgraph.graph.state import CompiledStateGraph
|
|
from langgraph.prebuilt import create_react_agent
|
|
from langgraph.types import interrupt, Command
|
|
from pydantic import BaseModel, Field
|
|
from starlette.responses import StreamingResponse
|
|
|
|
from app.model.ConversationModel import ConversationService
|
|
from app.utils.db_utils import AsyncSessionDep
|
|
from app.utils.llm_utils import create_llm
|
|
from app.utils.postgres_checkpointer import PostgresCheckpointerManager, AsyncPostgresSaverDep
|
|
|
|
|
|
# /*---------------------------------------获取时间工具-------------------------------------------*/
|
|
@tool(
|
|
name_or_callable="tool_get_datetime",
|
|
description="一个用于获取当前时间的工具,没有参数"
|
|
)
|
|
def tool_get_datetime():
|
|
return {
|
|
"result": f"当前时间为:{datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}",
|
|
"render": {
|
|
"type": "DataTable"
|
|
}
|
|
}
|
|
|
|
|
|
@tool(
|
|
name_or_callable="tool_multiply",
|
|
description="一个乘法工具,用于计算两个数字相乘"
|
|
)
|
|
def tool_multiply(
|
|
number1: Annotated[float, "第一个要相乘的数字"],
|
|
number2: Annotated[float, "第二个要相乘的数字"],
|
|
) -> float:
|
|
print("\n invoke tool_multiply", number1, number2, '\n')
|
|
|
|
return float(number1) * float(number2)
|
|
|
|
|
|
@tool(
|
|
name_or_callable="tool_add",
|
|
description="一个加法工具,用于计算两个数字相加"
|
|
)
|
|
def tool_add(
|
|
number1: Annotated[float, "第一个要相加的数字"],
|
|
number2: Annotated[float, "第二个要相加的数字"],
|
|
) -> float:
|
|
print("\n invoke tool_add", number1, number2, '\n')
|
|
|
|
return float(number1) + float(number2)
|
|
|
|
|
|
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
|
def tool_book_hotel() -> str:
|
|
resume_data = interrupt({
|
|
"title": "请确认酒店预定信息",
|
|
"form": [
|
|
{
|
|
"field": "hotel_name",
|
|
"type": "input",
|
|
"label": "酒店名称",
|
|
"required": True,
|
|
},
|
|
{
|
|
"field": "room_type",
|
|
"type": "select",
|
|
"label": "客房类型",
|
|
"options": [
|
|
{"label": "标间", "value": "标间"},
|
|
{"label": "单间", "value": "单间"},
|
|
{"label": "双人间", "value": "双人间"},
|
|
],
|
|
"required": True,
|
|
},
|
|
{
|
|
"field": "check_in_date",
|
|
"type": "date",
|
|
"label": "入住时间",
|
|
}
|
|
],
|
|
"formData": {
|
|
"hotel_name": '',
|
|
"room_type": '',
|
|
"check_in_date": '',
|
|
}
|
|
})
|
|
if resume_data == "N":
|
|
return f"用户选择取消预定酒店"
|
|
|
|
return ToolMessage(content="预定成功", additional_kwargs={"render": {"type": "DataTable"}})
|
|
|
|
|
|
class ChatMessage(BaseModel):
|
|
id: str = Field(..., description="消息id")
|
|
type: str = Field(..., description="消息类型")
|
|
content: str = Field(..., description="消息内容")
|
|
|
|
|
|
# 对话接口参数类型
|
|
class ChatParam(BaseModel):
|
|
thread_id: str = Field(..., description="线程id")
|
|
human_message: ChatMessage = Field(..., description="用户消息")
|
|
|
|
|
|
class ChatAgent:
|
|
agent: Union[CompiledStateGraph, None] = None
|
|
|
|
@staticmethod
|
|
async def get_agent() -> CompiledStateGraph:
|
|
if not ChatAgent.agent:
|
|
ChatAgent.agent = create_react_agent(
|
|
model=create_llm(),
|
|
tools=[
|
|
tool_book_hotel,
|
|
tool_get_datetime,
|
|
tool_add,
|
|
tool_multiply,
|
|
],
|
|
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
|
prompt="""
|
|
你是一名擅长使用工具的智能助手,你需要根据用户问题来进行回答,请使用中文进行回答。
|
|
当用户问题需要调用工具时再调用工具,否则按照你的知识来回答问题。
|
|
某些工具会触发中断让用户来编辑工具执行参数,这些工具会将新的执行参数作为信息返回,你需要回复用户最新的信息
|
|
"""
|
|
)
|
|
return ChatAgent.agent
|
|
|
|
@staticmethod
|
|
async def get_chat_state(thread_id: str):
|
|
graph = await ChatAgent.get_agent()
|
|
graph_state = await graph.aget_state(config={"configurable": {"thread_id": thread_id}})
|
|
if graph_state.values.get('messages', None) is None:
|
|
graph_state.values['messages'] = []
|
|
return {
|
|
**graph_state.values,
|
|
"__interrupt__": graph_state.interrupts,
|
|
}
|
|
|
|
|
|
def add_langgraph_chat_route(app: FastAPI):
|
|
# 流式对话接口
|
|
@app.post("/langgraph/stream")
|
|
async def langgraph_stream(body: dict, checkpointer: AsyncPostgresSaverDep):
|
|
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
|
print(body)
|
|
|
|
stream_input = body.get('input')
|
|
stream_config = body.get('config')
|
|
|
|
print("stream_input", stream_input)
|
|
print("stream_config", stream_config)
|
|
|
|
# 暂时每次对话的时候清理掉对话历史
|
|
# await checkpointer.adelete_thread(stream_config.get('configurable').get('thread_id'))
|
|
|
|
async def generator_function():
|
|
|
|
graph = await ChatAgent.get_agent()
|
|
# chat_state = await ChatAgent.get_chat_state(thread_id)
|
|
# chat_history_list = chat_state.get('messages')
|
|
async for chunk in graph.astream(
|
|
stream_input,
|
|
config=stream_config,
|
|
stream_mode=['messages', 'updates']
|
|
):
|
|
print("chunk-------->>>>>>>>>>", chunk)
|
|
|
|
result_template = {
|
|
"choices": [{"delta": {}, "index": 0}],
|
|
"created": time.time(),
|
|
"id": "",
|
|
"usage": None
|
|
}
|
|
|
|
emit_chunk = {"stream_type": chunk[0], }
|
|
|
|
if emit_chunk['stream_type'] == "messages":
|
|
# messages模式流式输出,此时 chunk[1][0] 为AIMessageChunk
|
|
chunk_message = chunk[1][0]
|
|
else:
|
|
# updates模式流式输出
|
|
for k, v in chunk[1].items():
|
|
if k == 'agent' or k == 'tool':
|
|
chunk_message = v.get('messages')[0]
|
|
elif k == '__interrupt__':
|
|
chunk_message = AIMessage(content='')
|
|
emit_chunk['stream_type'] = 'interrupt'
|
|
emit_chunk['interrupt'] = v[0].value
|
|
|
|
emit_chunk = {
|
|
**emit_chunk,
|
|
"msg_id": chunk_message.id,
|
|
"msg_type": chunk_message.type,
|
|
"msg_content": chunk_message.content,
|
|
}
|
|
|
|
if isinstance(chunk_message, AIMessage):
|
|
if chunk_message.tool_calls:
|
|
emit_chunk["tool_calls"] = chunk_message.tool_calls
|
|
elif isinstance(chunk_message, ToolMessage):
|
|
emit_chunk["tool_name"] = chunk_message.name
|
|
emit_chunk["additional_kwargs"] = chunk_message.additional_kwargs
|
|
|
|
result_template["choices"][0]["delta"]['content'] = emit_chunk
|
|
result_template["choices"][0]["delta"]['role'] = 'assistant'
|
|
|
|
result_template['id'] = chunk_message.id
|
|
result_template['created'] = int(time.time())
|
|
|
|
yield f"data: {json.dumps(result_template, ensure_ascii=False)}\n\n"
|
|
# yield "data: [DONE]\n\n"
|
|
|
|
return StreamingResponse(generator_function(), media_type="text/event-stream")
|
|
|
|
# 恢复中断接口
|
|
@app.post("/langgraph/chat_resume/{thread_id}")
|
|
async def langgraph_chat(body: dict, thread_id: str):
|
|
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}}
|
|
)
|
|
return {
|
|
**graph_state,
|
|
# 这里不需要加1,因为我们并没有往messages中增加消息
|
|
"messages": graph_state.get('messages')[len(chat_history_list):],
|
|
}
|
|
|
|
# 查询聊天记录
|
|
@app.get("/langgraph/chat_state/{thread_id}")
|
|
async def langgraph_chat(thread_id: str):
|
|
return await ChatAgent.get_chat_state(thread_id) # 查询聊天记录
|
|
|
|
# 删除聊天记录
|
|
@app.post("/langgraph/chat_remove/{thread_id}")
|
|
async def langgraph_chat(thread_id: str, checkpointer: AsyncPostgresSaverDep, session: AsyncSessionDep):
|
|
await ConversationService.item_delete(session=session, row_dict={"id": thread_id})
|
|
await checkpointer.adelete_thread(thread_id)
|
|
return {"result": "success"}
|