feat: 准备实现工具的中断恢复功能

This commit is contained in:
martsforever
2025-09-04 13:47:14 +08:00
parent dcce07fc0d
commit 9055a56224
+30 -30
View File
@@ -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():
chunk_message = v.get('messages')[0] 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 = {
**emit_chunk, **emit_chunk,