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

@@ -4,3 +4,8 @@ class ChatRequest(BaseModel):
username: str username: str
message: str message: str
thread_id: 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 import APIRouter
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from model import ChatRequest from model import ChatRequest, SessionMessageResponse
from services.agent_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()
@@ -16,11 +18,16 @@ async def chat_endpoint(request: ChatRequest):
) )
@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

@@ -2,7 +2,7 @@ from langchain_core.messages import HumanMessage, AIMessageChunk, AIMessage
from agent import agent, checkpointer, title_agent from agent import agent, checkpointer, title_agent
from model import ChatRequest 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 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}) 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)

View File

@@ -1,5 +1,8 @@
import sqlite3 import sqlite3
import os import os
from typing import List
from model import SessionMessageResponse
def init_db(): def init_db():
@@ -20,6 +23,34 @@ def init_db():
conn.close() 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: def exists_session(thread_id: str, username: str) -> bool:
conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH")) conn = sqlite3.connect(os.getenv("SQLITE_DB_PATH"))
cursor = conn.cursor() cursor = conn.cursor()