from langchain_core.messages import HumanMessage, AIMessageChunk, AIMessage from agent import agent, checkpointer, title_agent from model import ChatRequest from services.db_service import exists_session, save_session, query_session_messages from utils import dicts_to_messages def generate_title(messages: list[dict[str, str]]) -> str: response = title_agent.invoke({ "messages": dicts_to_messages(messages) }) return response["messages"][-1].content.strip() async def search_recipes(request: ChatRequest): """调用agent搜索食谱""" print(f"[用户] {request.username}: {request.message}, thread_id: {request.thread_id}") try: message = HumanMessage(content=request.message) # 流式调用Agent for chunk, metadata in agent.stream( {"messages": [message]}, {"configurable": {"thread_id": request.thread_id}}, stream_mode="messages" ): if isinstance(chunk, AIMessageChunk) and chunk.content: yield chunk.content # 总结对话标题并保存 if not exists_session(request.thread_id, request.username): title = generate_title(get_messages(request.thread_id)) save_session(request.thread_id, request.username, title) except Exception as e: print(f"\n[错误]: {str(e)}") yield "信息检索失败,试试看手动输入食物列表?" # 清空会话 def clear_messages(thread_id: str): """清空会话""" print(f"清空历史消息,thread_id: {thread_id}") checkpointer.delete_thread(thread_id) # 查询会话历史 def get_messages(thread_id: str) -> list[dict[str, str]]: """获取会话历史""" print(f"获取历史消息,thread_id: {thread_id}") # 根据 thread_id 查询 checkpoint checkpoint = checkpointer.get({"configurable": {"thread_id": thread_id}}) # 如果不存在,返回空列表 if not checkpoint: return [] # 安全获取 messages channel_values = checkpoint.get("channel_values") if not channel_values: return [] messages = channel_values.get("messages", []) if not messages: return [] # 转换消息格式 result = [] for msg in messages: if not msg.content: continue if isinstance(msg, HumanMessage): result.append({"role": "user", "content": msg.content}) elif isinstance(msg, AIMessage): result.append({"role": "ai", "content": msg.content}) return result # 查询历史会话消息 def get_user_session_messages(username: str): """查询历史会话消息""" print(f"查询历史会话消息,username: {username}") return query_session_messages(username)