feat:增加查询用户历史消息功能
This commit is contained in:
BIN
chief.db-shm
Normal file
BIN
chief.db-shm
Normal file
Binary file not shown.
0
chief.db-wal
Normal file
0
chief.db-wal
Normal file
5
model.py
5
model.py
@@ -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
|
||||||
19
router.py
19
router.py
@@ -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")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user