feat: 准备实现工具的中断恢复功能
This commit is contained in:
@@ -13,47 +13,46 @@ from pydantic import BaseModel, Field
|
|||||||
from starlette.responses import StreamingResponse
|
from starlette.responses import StreamingResponse
|
||||||
|
|
||||||
from app.utils.llm_utils import create_llm
|
from app.utils.llm_utils import create_llm
|
||||||
from app.utils.postgres_checkpointer import PostgresCheckpointerManager
|
from app.utils.postgres_checkpointer import PostgresCheckpointerManager, AsyncPostgresSaverDep
|
||||||
|
|
||||||
|
|
||||||
# /*---------------------------------------获取时间工具-------------------------------------------*/
|
# /*---------------------------------------获取时间工具-------------------------------------------*/
|
||||||
@tool(name_or_callable="获取时间工具", description="一个用于获取当前时间的工具,没有参数")
|
@tool(
|
||||||
|
name_or_callable="tool_get_datetime",
|
||||||
|
description="一个用于获取当前时间的工具,没有参数"
|
||||||
|
)
|
||||||
def tool_get_datetime():
|
def tool_get_datetime():
|
||||||
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||||
|
|
||||||
|
|
||||||
@tool(
|
@tool(
|
||||||
name_or_callable="乘法工具",
|
name_or_callable="tool_multiply",
|
||||||
description="一个乘法工具,用于计算两个数字相乘"
|
description="一个乘法工具,用于计算两个数字相乘"
|
||||||
)
|
)
|
||||||
def multiple_tool(
|
def tool_multiply(
|
||||||
number1: Annotated[float, "第一个要相乘的数字"],
|
number1: Annotated[float, "第一个要相乘的数字"],
|
||||||
number2: Annotated[float, "第二个要相乘的数字"],
|
number2: Annotated[float, "第二个要相乘的数字"],
|
||||||
) -> float:
|
) -> float:
|
||||||
print("\n invoke multiple_tool", number1, number2, '\n')
|
print("\n invoke tool_multiply", number1, number2, '\n')
|
||||||
|
|
||||||
return float(number1) * float(number2)
|
return float(number1) * float(number2)
|
||||||
|
|
||||||
|
|
||||||
@tool(
|
@tool(
|
||||||
name_or_callable="加法工具",
|
name_or_callable="tool_add",
|
||||||
description="一个加法工具,用于计算两个数字相加"
|
description="一个加法工具,用于计算两个数字相加"
|
||||||
)
|
)
|
||||||
def add_tool(
|
def tool_add(
|
||||||
number1: Annotated[float, "第一个要相加的数字"],
|
number1: Annotated[float, "第一个要相加的数字"],
|
||||||
number2: Annotated[float, "第二个要相加的数字"],
|
number2: Annotated[float, "第二个要相加的数字"],
|
||||||
) -> float:
|
) -> float:
|
||||||
print("\n invoke add_tool", number1, number2, '\n')
|
print("\n invoke tool_add", number1, number2, '\n')
|
||||||
|
|
||||||
return float(number1) + float(number2)
|
return float(number1) + float(number2)
|
||||||
|
|
||||||
|
|
||||||
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
||||||
def tool_book_hotel(
|
def tool_book_hotel() -> str:
|
||||||
hotel_name: Annotated[str, '酒店名称'],
|
|
||||||
room_type: Annotated[str, '房间类型'],
|
|
||||||
check_in_date: Annotated[str, '入住时间'],
|
|
||||||
) -> str:
|
|
||||||
resume_data = interrupt({
|
resume_data = interrupt({
|
||||||
"title": "请确认酒店预定信息",
|
"title": "请确认酒店预定信息",
|
||||||
"form": [
|
"form": [
|
||||||
@@ -81,24 +80,15 @@ def tool_book_hotel(
|
|||||||
}
|
}
|
||||||
],
|
],
|
||||||
"formData": {
|
"formData": {
|
||||||
"hotel_name": hotel_name,
|
"hotel_name": '',
|
||||||
"room_type": room_type,
|
"room_type": '',
|
||||||
"check_in_date": check_in_date,
|
"check_in_date": '',
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
if resume_data == "N":
|
if resume_data == "N":
|
||||||
return f"用户选择取消预定酒店"
|
return f"用户选择取消预定酒店"
|
||||||
|
|
||||||
old_book_data = {
|
return "预定成功"
|
||||||
"hotel_name": hotel_name,
|
|
||||||
"room_type": room_type,
|
|
||||||
"check_in_date": check_in_date,
|
|
||||||
}
|
|
||||||
new_book_data: dict = resume_data
|
|
||||||
|
|
||||||
change_field_list = [(k, v) for k, v in old_book_data.items() if old_book_data.get(k) != new_book_data.get(k)]
|
|
||||||
|
|
||||||
return (f"用户已经更改参数,最新的为{json.dumps(new_book_data, ensure_ascii=False)}" if change_field_list else "") + ",结果为预定成功"
|
|
||||||
|
|
||||||
|
|
||||||
class ChatMessage(BaseModel):
|
class ChatMessage(BaseModel):
|
||||||
@@ -124,8 +114,8 @@ class ChatAgent:
|
|||||||
tools=[
|
tools=[
|
||||||
tool_book_hotel,
|
tool_book_hotel,
|
||||||
tool_get_datetime,
|
tool_get_datetime,
|
||||||
add_tool,
|
tool_add,
|
||||||
multiple_tool,
|
tool_multiply,
|
||||||
],
|
],
|
||||||
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
||||||
prompt="""
|
prompt="""
|
||||||
@@ -151,7 +141,7 @@ 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):
|
async def langgraph_stream(body: dict, checkpointer: AsyncPostgresSaverDep):
|
||||||
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
||||||
print(body)
|
print(body)
|
||||||
|
|
||||||
@@ -161,6 +151,9 @@ def add_langgraph_chat_route(app: FastAPI):
|
|||||||
print("stream_input", stream_input)
|
print("stream_input", stream_input)
|
||||||
print("stream_config", stream_config)
|
print("stream_config", stream_config)
|
||||||
|
|
||||||
|
# 暂时每次对话的时候清理掉对话历史
|
||||||
|
await checkpointer.adelete_thread(stream_config.get('configurable').get('thread_id'))
|
||||||
|
|
||||||
async def generator_function():
|
async def generator_function():
|
||||||
|
|
||||||
graph = await ChatAgent.get_agent()
|
graph = await ChatAgent.get_agent()
|
||||||
@@ -183,10 +176,17 @@ def add_langgraph_chat_route(app: FastAPI):
|
|||||||
emit_chunk = {"stream_type": chunk[0], }
|
emit_chunk = {"stream_type": chunk[0], }
|
||||||
|
|
||||||
if emit_chunk['stream_type'] == "messages":
|
if emit_chunk['stream_type'] == "messages":
|
||||||
|
# messages模式流式输出,此时 chunk[1][0] 为AIMessageChunk
|
||||||
chunk_message = chunk[1][0]
|
chunk_message = chunk[1][0]
|
||||||
else:
|
else:
|
||||||
|
# updates模式流式输出
|
||||||
for k, v in chunk[1].items():
|
for k, v in chunk[1].items():
|
||||||
|
if k == 'agent' or k == 'tool':
|
||||||
chunk_message = v.get('messages')[0]
|
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 = {
|
||||||
**emit_chunk,
|
**emit_chunk,
|
||||||
|
|||||||
Reference in New Issue
Block a user