From 9055a56224419254179123a267228088770ef8a7 Mon Sep 17 00:00:00 2001 From: martsforever Date: Thu, 4 Sep 2025 13:47:14 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=87=86=E5=A4=87=E5=AE=9E=E7=8E=B0?= =?UTF-8?q?=E5=B7=A5=E5=85=B7=E7=9A=84=E4=B8=AD=E6=96=AD=E6=81=A2=E5=A4=8D?= =?UTF-8?q?=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/controller/add_langgraph_chat_route.py | 60 +++++++++++----------- 1 file changed, 30 insertions(+), 30 deletions(-) diff --git a/app/controller/add_langgraph_chat_route.py b/app/controller/add_langgraph_chat_route.py index 2400563..bcf2978 100644 --- a/app/controller/add_langgraph_chat_route.py +++ b/app/controller/add_langgraph_chat_route.py @@ -13,47 +13,46 @@ from pydantic import BaseModel, Field from starlette.responses import StreamingResponse 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(): return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") @tool( - name_or_callable="乘法工具", + name_or_callable="tool_multiply", description="一个乘法工具,用于计算两个数字相乘" ) -def multiple_tool( +def tool_multiply( number1: Annotated[float, "第一个要相乘的数字"], number2: Annotated[float, "第二个要相乘的数字"], ) -> float: - print("\n invoke multiple_tool", number1, number2, '\n') + print("\n invoke tool_multiply", number1, number2, '\n') return float(number1) * float(number2) @tool( - name_or_callable="加法工具", + name_or_callable="tool_add", description="一个加法工具,用于计算两个数字相加" ) -def add_tool( +def tool_add( number1: Annotated[float, "第一个要相加的数字"], number2: Annotated[float, "第二个要相加的数字"], ) -> float: - print("\n invoke add_tool", number1, number2, '\n') + 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( - hotel_name: Annotated[str, '酒店名称'], - room_type: Annotated[str, '房间类型'], - check_in_date: Annotated[str, '入住时间'], -) -> str: +def tool_book_hotel() -> str: resume_data = interrupt({ "title": "请确认酒店预定信息", "form": [ @@ -81,24 +80,15 @@ def tool_book_hotel( } ], "formData": { - "hotel_name": hotel_name, - "room_type": room_type, - "check_in_date": check_in_date, + "hotel_name": '', + "room_type": '', + "check_in_date": '', } }) if resume_data == "N": return f"用户选择取消预定酒店" - old_book_data = { - "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 "") + ",结果为预定成功" + return "预定成功" class ChatMessage(BaseModel): @@ -124,8 +114,8 @@ class ChatAgent: tools=[ tool_book_hotel, tool_get_datetime, - add_tool, - multiple_tool, + tool_add, + tool_multiply, ], checkpointer=await PostgresCheckpointerManager.get_instance(), prompt=""" @@ -151,7 +141,7 @@ class ChatAgent: def add_langgraph_chat_route(app: FastAPI): # 流式对话接口 @app.post("/langgraph/stream") - async def langgraph_stream(body: dict): + async def langgraph_stream(body: dict, checkpointer: AsyncPostgresSaverDep): print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::") print(body) @@ -161,6 +151,9 @@ def add_langgraph_chat_route(app: FastAPI): 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() @@ -183,10 +176,17 @@ def add_langgraph_chat_route(app: FastAPI): 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(): - 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,