import datetime import json import time from typing import Union, Annotated, TypedDict import httpx from fastapi import FastAPI, Depends from langchain_core.messages import AIMessage, ToolMessage from langchain_core.runnables import RunnableConfig 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.config.env import env from app.controller.add_user_route import oauth2_scheme 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 f"当前时间为:{datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}" @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) class BookHotelData(TypedDict): hotel_id: str proj_id: str user_id: str @tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具") async def tool_book_hotel(config: RunnableConfig) -> str: """ 触发中断,输入酒店预定表单信息 """ book_hotel_data: Union[BookHotelData, str] = interrupt({"title": "请输入酒店预定信息", "formCode": "bookHotel"}) if book_hotel_data == "N": return f"用户选择取消预定酒店" hotel_id = book_hotel_data.get('hotel_id') proj_id = book_hotel_data.get('proj_id') user_id = book_hotel_data.get('user_id') token = config['configurable']['token'] # print("================ tool_book_hotel ================") # print("hotel_id", hotel_id) # print("proj_id", proj_id) # print("user_id", user_id) async with httpx.AsyncClient() as client: response = await client.post( url=f"http://localhost:{env.server_port}/book_hotel", headers={ "Authorization": f"Bearer {token}" }, json={ "hotel_id": hotel_id, "user_id": user_id, "proj_id": proj_id, }, ) response_data = response.json() # print("================ response_data ================") # print(response_data) approve_id = response_data.get('graph_state').get('input_approve_id') return [ "预定成功," + response_data.get('message'), { "component": "Link", "props": { "to": f"/pages/approve/approve-detail?id={approve_id}", "children": "点击查看审批进度" } } ] 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("huoshan-think-pro"), 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, token: str = Depends(oauth2_scheme), ): print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::") print(body) stream_input = body.get('input') stream_config = body.get('config') stream_config['configurable']['token'] = token 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, token: str = Depends(oauth2_scheme), ): 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, "token": token}} ) 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"}