fix: paginate chat detail messages
This commit is contained in:
@@ -83,6 +83,8 @@ def export_chats(
|
||||
@router.get("/chat/{session_id}")
|
||||
def chat_detail(
|
||||
session_id: int,
|
||||
messagePage: int = Query(default=1, ge=1),
|
||||
messagePageSize: int = Query(default=20, ge=10, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
@@ -97,10 +99,13 @@ def chat_detail(
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="会话不存在")
|
||||
|
||||
session, user = session_row
|
||||
message_query = select(ChatMessage).where(ChatMessage.session_id == session_id)
|
||||
message_total = db.scalar(select(func.count()).select_from(message_query.subquery())) or 0
|
||||
messages = db.scalars(
|
||||
select(ChatMessage)
|
||||
.where(ChatMessage.session_id == session_id)
|
||||
message_query
|
||||
.order_by(ChatMessage.created_at.asc(), ChatMessage.id.asc())
|
||||
.offset((messagePage - 1) * messagePageSize)
|
||||
.limit(messagePageSize)
|
||||
).all()
|
||||
ai_logs = db.scalars(
|
||||
select(AiRequestLog)
|
||||
@@ -111,6 +116,12 @@ def chat_detail(
|
||||
{
|
||||
"session": _chat_row_dict(session, user),
|
||||
"messages": [_message_dict(item) for item in messages],
|
||||
"messagesPage": page_result(
|
||||
[_message_dict(item) for item in messages],
|
||||
total=message_total,
|
||||
page=messagePage,
|
||||
page_size=messagePageSize,
|
||||
),
|
||||
"aiLogs": [_ai_log_dict(item, include_prompt=True) for item in ai_logs],
|
||||
}
|
||||
)
|
||||
|
||||
@@ -3,9 +3,10 @@ from sqlalchemy.orm import Session
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from app.api.admin_agent_records import attention_list, retrieval_logs
|
||||
from app.api.admin_records import ai_logs
|
||||
from app.api.admin_records import ai_logs, chat_detail
|
||||
from app.api.admin_users import list_users
|
||||
from app.models import Base
|
||||
from app.models.chat import ChatMessage, ChatSession
|
||||
from app.models.knowledge import HumanAttentionRecord, KnowledgeRetrievalLog
|
||||
from app.models.logs import AiRequestLog
|
||||
from app.models.user import User
|
||||
@@ -43,6 +44,27 @@ def test_ai_log_page_does_not_load_large_detail_fields():
|
||||
assert all(item["retrievedChunks"] == [] for item in response["data"]["items"])
|
||||
|
||||
|
||||
def test_chat_detail_messages_are_paginated():
|
||||
with _database() as db:
|
||||
user = User(id=1, phone="13800000000", name="学员", daily_chat_limit=10)
|
||||
session = ChatSession(id=1, user_id=1, title="长会话", message_count=25)
|
||||
db.add_all([user, session])
|
||||
db.add_all([
|
||||
ChatMessage(id=index + 1, session_id=1, user_id=1, role="user", content=f"消息{index + 1}")
|
||||
for index in range(25)
|
||||
])
|
||||
db.commit()
|
||||
|
||||
response = chat_detail(1, messagePage=2, messagePageSize=10, db=db, current_admin=object())
|
||||
|
||||
data = response["data"]
|
||||
assert data["session"]["messageCount"] == 25
|
||||
assert data["messagesPage"]["total"] == 25
|
||||
assert data["messagesPage"]["page"] == 2
|
||||
assert len(data["messages"]) == 10
|
||||
assert data["messages"][0]["content"] == "消息11"
|
||||
|
||||
|
||||
def test_retrieval_and_attention_lists_are_paginated():
|
||||
with _database() as db:
|
||||
for index in range(21):
|
||||
|
||||
Reference in New Issue
Block a user