Compare commits

...

5 Commits

Author SHA1 Message Date
54eb3680f1 feat:删除数据库文件 2026-05-31 21:22:40 +08:00
b04242bbb2 feat:增加查询用户历史消息功能 2026-05-31 21:20:25 +08:00
c45ce61e9c feat:删除数据库文件 2026-05-28 14:07:24 +08:00
e0b7ed4364 feat:合并冲突 2026-05-28 14:06:02 +08:00
c136426589 feat:增加存储用户对话记录 2026-05-28 13:55:37 +08:00
12 changed files with 180 additions and 23 deletions

3
.gitignore vendored
View File

@@ -66,3 +66,6 @@ media/
logs/ logs/
packages/ packages/
*.db
*.db-shm
*.db-wal

View File

@@ -5,16 +5,24 @@ import sqlite3
from langchain_tavily import TavilySearch from langchain_tavily import TavilySearch
from langgraph.checkpoint.sqlite import SqliteSaver from langgraph.checkpoint.sqlite import SqliteSaver
from services.db_service import init_db
from dotenv import load_dotenv
import os
load_dotenv()
# 连接sqlite # 连接sqlite
connection = sqlite3.connect("chief.db", check_same_thread=False) connection = sqlite3.connect(os.getenv("SQLITE_DB_PATH"), check_same_thread=False)
# 初始化checkpointer # 初始化checkpointer
checkpointer = SqliteSaver(connection) checkpointer = SqliteSaver(connection)
# 自动建表 # 自动建表
checkpointer.setup() checkpointer.setup()
init_db()
# web搜索工具使用tavily作为web搜索工具 # web搜索工具使用tavily作为web搜索工具
web_search = TavilySearch( web_search = TavilySearch(
tavily_api_key="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", tavily_api_key=os.getenv("TAVILY_API_KEY"),
max_results=5, max_results=5,
topic="general" topic="general"
) )
@@ -30,10 +38,10 @@ system_prompt = """
""" """
model = init_chat_model( model = init_chat_model(
model="qwen-turbo", model=os.getenv("MODEL_NAME"),
model_provider="openai", model_provider="openai",
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1", base_url=os.getenv("MODEL_BASE_URL"),
api_key="sk-52bcd98e9c1d45908437c4e8706eefff", api_key=os.getenv("MODEL_API_KEY"),
temperature=1.5, temperature=1.5,
) )
@@ -43,3 +51,13 @@ agent = create_agent(
checkpointer=checkpointer, # 记忆 checkpointer=checkpointer, # 记忆
system_prompt=system_prompt # 系统提示词 system_prompt=system_prompt # 系统提示词
) )
title_prompt = """
请根据以下对话内容,生成一句不超过 20 字的标题,用于记录本次对话主题。
只返回标题,不要解释。
"""
title_agent = create_agent(
model=model,
system_prompt=title_prompt
)

BIN
chief.db

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@@ -1,4 +1,11 @@
from pydantic import BaseModel from pydantic import BaseModel
class ChatRequest(BaseModel): class ChatRequest(BaseModel):
username: str
message: str message: str
thread_id: str
class SessionMessageResponse(BaseModel):
username: str
title: str
thread_id: str

View File

@@ -5,5 +5,6 @@ langchain-openai~=1.2.2
langchain-tavily~=0.2.18 langchain-tavily~=0.2.18
langgraph~=1.2.2 langgraph~=1.2.2
langgraph-checkpoint-sqlite~=3.1.0 langgraph-checkpoint-sqlite~=3.1.0
pydantic~=2.12.5 pydantic~=2.13.4
langchain-core~=1.4.0 langchain-core~=1.4.0
python-dotenv~=1.2.2

View File

@@ -1,10 +1,10 @@
import uuid from typing import List
from fastapi import APIRouter from fastapi import APIRouter
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from model import ChatRequest from model import ChatRequest, SessionMessageResponse
from service import get_messages, clear_messages, search_recipes from services.agent_service import get_messages, clear_messages, search_recipes, get_user_session_messages
router = APIRouter() router = APIRouter()
@@ -13,16 +13,21 @@ router = APIRouter()
async def chat_endpoint(request: ChatRequest): async def chat_endpoint(request: ChatRequest):
"""流式对话""" """流式对话"""
return StreamingResponse( return StreamingResponse(
search_recipes(request.message, str(uuid.uuid4())), search_recipes(request),
media_type="text/event-stream" media_type="text/event-stream"
) )
@router.get("/chat/messages") @router.get("/chat/messages/session", response_model=List[SessionMessageResponse])
async def get_chat_messages(thread_id: str): def get_session_messages(username: str):
"""获取历史消息""" """获取历史消息"""
messages = get_messages(thread_id) return get_user_session_messages(username)
return {"messages": messages}
@router.get("/chat/messages")
def get_chat_messages(thread_id: str):
"""获取历史消息"""
return get_messages(thread_id)
@router.delete("/chat/messages") @router.delete("/chat/messages")

View File

@@ -1,23 +1,38 @@
from langchain_core.messages import HumanMessage, AIMessageChunk, AIMessage from langchain_core.messages import HumanMessage, AIMessageChunk, AIMessage
from agent import agent, checkpointer 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
async def search_recipes(prompt: str, thread_id: str): 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搜索食谱""" """调用agent搜索食谱"""
print(f"[用户]: {prompt}, thread_id: {thread_id}") print(f"[用户] {request.username}: {request.message}, thread_id: {request.thread_id}")
try: try:
message = HumanMessage(content=prompt) message = HumanMessage(content=request.message)
# 流式调用Agent # 流式调用Agent
for chunk, metadata in agent.stream( for chunk, metadata in agent.stream(
{"messages": [message]}, {"messages": [message]},
{"configurable": {"thread_id": thread_id}}, {"configurable": {"thread_id": request.thread_id}},
stream_mode="messages" stream_mode="messages"
): ):
if isinstance(chunk, AIMessageChunk) and chunk.content: if isinstance(chunk, AIMessageChunk) and chunk.content:
yield 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: except Exception as e:
print(f"\n[错误]: {str(e)}") print(f"\n[错误]: {str(e)}")
yield "信息检索失败,试试看手动输入食物列表?" yield "信息检索失败,试试看手动输入食物列表?"
@@ -60,6 +75,13 @@ def get_messages(thread_id: str) -> list[dict[str, str]]:
if isinstance(msg, HumanMessage): if isinstance(msg, HumanMessage):
result.append({"role": "user", "content": msg.content}) result.append({"role": "user", "content": msg.content})
elif isinstance(msg, AIMessage): elif isinstance(msg, AIMessage):
result.append({"role": "assistant", "content": msg.content}) result.append({"role": "ai", "content": msg.content})
return result return result
# 查询历史会话消息
def get_user_session_messages(username: str):
"""查询历史会话消息"""
print(f"查询历史会话消息username: {username}")
return query_session_messages(username)

85
services/db_service.py Normal file
View File

@@ -0,0 +1,85 @@
import sqlite3
import os
from typing import List
from model import SessionMessageResponse
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 query_session_messages(username: str) -> List[SessionMessageResponse]:
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
cursor.execute(
"""
SELECT username, title, thread_id
FROM session
WHERE username = ?
ORDER BY created_time DESC
""",
(username,)
)
rows = cursor.fetchall()
conn.close()
return [
SessionMessageResponse(
username=row["username"],
title=row["title"],
thread_id=row["thread_id"],
)
for row in rows
]
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()

5
test.py Normal file
View File

@@ -0,0 +1,5 @@
from agent import model
if __name__ == '__main__':
response = model.invoke("西红柿、鸡蛋")
print(response)

11
utils.py Normal file
View File

@@ -0,0 +1,11 @@
from langchain_core.messages import HumanMessage, AIMessage
def dicts_to_messages(messages: list[dict]):
result = []
for m in messages:
if m["role"] == "user":
result.append(HumanMessage(content=m["content"]))
elif m["role"] == "ai":
result.append(AIMessage(content=m["content"]))
return result