feat:增加查询用户历史消息功能

This commit is contained in:
2026-05-31 21:20:25 +08:00
parent c45ce61e9c
commit b04242bbb2
6 changed files with 57 additions and 7 deletions

BIN
chief.db-shm Normal file

Binary file not shown.

0
chief.db-wal Normal file
View File

View File

@@ -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

View File

@@ -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")

View File

@@ -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)

View File

@@ -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()