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"}