Files
chief-agent/services/agent_service.py

88 lines
2.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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)