diff --git a/chief.db-shm b/chief.db-shm new file mode 100644 index 0000000..fe9ac28 Binary files /dev/null and b/chief.db-shm differ diff --git a/chief.db-wal b/chief.db-wal new file mode 100644 index 0000000..e69de29 diff --git a/model.py b/model.py index 15f309a..22d07b8 100644 --- a/model.py +++ b/model.py @@ -3,4 +3,9 @@ from pydantic import BaseModel class ChatRequest(BaseModel): username: str message: str + thread_id: str + +class SessionMessageResponse(BaseModel): + username: str + title: str thread_id: str \ No newline at end of file diff --git a/router.py b/router.py index 84205be..624dd33 100644 --- a/router.py +++ b/router.py @@ -1,8 +1,10 @@ +from typing import List + from fastapi import APIRouter from fastapi.responses import StreamingResponse -from model import ChatRequest -from services.agent_service import get_messages, clear_messages, search_recipes +from model import ChatRequest, SessionMessageResponse +from services.agent_service import get_messages, clear_messages, search_recipes, get_user_session_messages router = APIRouter() @@ -16,11 +18,16 @@ async def chat_endpoint(request: ChatRequest): ) -@router.get("/chat/messages") -async def get_chat_messages(thread_id: str): +@router.get("/chat/messages/session", response_model=List[SessionMessageResponse]) +def get_session_messages(username: str): """获取历史消息""" - messages = get_messages(thread_id) - return {"messages": messages} + return get_user_session_messages(username) + + +@router.get("/chat/messages") +def get_chat_messages(thread_id: str): + """获取历史消息""" + return get_messages(thread_id) @router.delete("/chat/messages") diff --git a/services/agent_service.py b/services/agent_service.py index 2b38bac..e0650bd 100644 --- a/services/agent_service.py +++ b/services/agent_service.py @@ -2,7 +2,7 @@ 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 services.db_service import exists_session, save_session, query_session_messages from utils import dicts_to_messages @@ -78,3 +78,10 @@ def get_messages(thread_id: str) -> list[dict[str, str]]: 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) diff --git a/services/db_service.py b/services/db_service.py index 9accfbd..e6497a6 100644 --- a/services/db_service.py +++ b/services/db_service.py @@ -1,5 +1,8 @@ import sqlite3 import os +from typing import List + +from model import SessionMessageResponse def init_db(): @@ -20,6 +23,34 @@ def init_db(): 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()