feat:增加存储用户对话记录

This commit is contained in:
2026-05-28 13:55:37 +08:00
parent 613f426aee
commit c136426589
11 changed files with 122 additions and 18 deletions

80
services/agent_service.py Normal file
View File

@@ -0,0 +1,80 @@
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
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

54
services/db_service.py Normal file
View File

@@ -0,0 +1,54 @@
import sqlite3
import os
def init_db():
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
# 业务表:会话
conn.execute("""
CREATE TABLE IF NOT EXISTS session (
id INTEGER PRIMARY KEY AUTOINCREMENT,
thread_id TEXT NOT NULL UNIQUE,
username TEXT NOT NULL,
title TEXT NOT NULL,
created_time DATETIME DEFAULT CURRENT_TIMESTAMP
);
""")
conn.commit()
conn.close()
def exists_session(thread_id: str, username: str) -> bool:
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
cursor = conn.cursor()
cursor.execute(
"""
SELECT 1
FROM session
WHERE thread_id = ? AND username = ?
LIMIT 1
""",
(thread_id, username)
)
result = cursor.fetchone()
conn.close()
return result is not None
def save_session(thread_id: str, username: str, title: str):
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
conn.execute(
"""
INSERT OR REPLACE INTO session
(thread_id, username, title)
VALUES (?, ?, ?)
""",
(thread_id, username, title)
)
conn.commit()
conn.close()