feat: improve chat feedback navigation

This commit is contained in:
2026-08-31 11:41:56 +08:00
parent 5740bd26e7
commit f4b2836c8a
11 changed files with 422 additions and 41 deletions

View File

@@ -0,0 +1,25 @@
"""add chat message navigation index
Revision ID: 0043_chat_msg_nav_index
Revises: 0042_record_export_permission
"""
from alembic import op
revision = "0043_chat_msg_nav_index"
down_revision = "0042_record_export_permission"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_index(
"ix_chat_message_session_created_id",
"sys_chat_message",
["session_id", "created_at", "id"],
)
def downgrade() -> None:
op.drop_index("ix_chat_message_session_created_id", table_name="sys_chat_message")

View File

@@ -29,6 +29,7 @@ from app.services.growth_profile_service import topic_dict, topic_summary_dict
from app.services.help_card_service import help_card_dict
from app.services.share_draft_service import share_draft_dict
from app.services.chat_export_service import ChatExportService
from app.services.chat_message_navigation_service import ChatMessageNavigationService
router = APIRouter()
@@ -119,6 +120,7 @@ def export_chats(
@router.get("/chat/{session_id}/messages")
def chat_messages(
session_id: int,
keyword: str = Query(default="", max_length=100),
page: int = Query(default=1, ge=1),
pageSize: int = Query(default=20, ge=10, le=100),
db: Session = Depends(get_db),
@@ -127,7 +129,10 @@ def chat_messages(
exists_session = db.scalar(select(ChatSession.id).where(ChatSession.id == session_id))
if exists_session is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="会话不存在")
normalized_keyword = keyword.strip()
message_query = select(ChatMessage).where(ChatMessage.session_id == session_id)
if normalized_keyword:
message_query = message_query.where(ChatMessage.content.contains(normalized_keyword, autoescape=True))
total = db.scalar(select(func.count()).select_from(message_query.subquery())) or 0
messages = db.scalars(
message_query
@@ -211,30 +216,18 @@ def chat_detail(
)
if focused_message is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="触发消息不属于该会话或已不存在")
focused_position = db.scalar(
select(func.count()).select_from(ChatMessage).where(
ChatMessage.session_id == session_id,
or_(
ChatMessage.created_at < focused_message.created_at,
(
(ChatMessage.created_at == focused_message.created_at)
& (ChatMessage.id <= focused_message.id)
),
),
)
) or 1
messagePage = (focused_position - 1) // messagePageSize + 1
messagePage = ChatMessageNavigationService.page_for_message(
db,
session_id=session_id,
message=focused_message,
page_size=messagePageSize,
)
messages = db.scalars(
message_query
.order_by(ChatMessage.created_at.asc(), ChatMessage.id.asc())
.offset((messagePage - 1) * messagePageSize)
.limit(messagePageSize)
).all()
ai_logs = db.scalars(
select(AiRequestLog)
.where(AiRequestLog.session_id == session_id)
.order_by(AiRequestLog.created_at.asc(), AiRequestLog.id.asc())
).all()
topics = _topic_rows(db, session_id)
topic_ids = [item["id"] for item in topics]
help_cards = []
@@ -264,7 +257,8 @@ def chat_detail(
page=messagePage,
page_size=messagePageSize,
),
"aiLogs": [_ai_log_dict(item, include_prompt=True) for item in ai_logs],
# Large prompts and retrieval chunks are loaded asynchronously by the drawer.
"aiLogs": [],
"topics": topics,
"helpCards": [help_card_dict(item) for item in help_cards],
"shareDrafts": [share_draft_dict(item) for item in share_drafts],

View File

@@ -1,7 +1,8 @@
from __future__ import annotations
from datetime import UTC, date, datetime, timedelta
from datetime import date, timedelta
from io import BytesIO
from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi.responses import StreamingResponse
@@ -18,13 +19,14 @@ from app.api.pagination import page_result
from app.core.database import get_db
from app.core.dependencies import get_current_admin, get_current_user
from app.core.responses import api_success
from app.core.time_utils import business_day_boundary, to_business_naive
from app.core.time_utils import business_day_boundary, to_business_naive, utc_now_naive
from app.models.admin import Admin
from app.models.chat import ChatMessage, ChatSession
from app.models.feedback import MessageFeedback
from app.models.user import User
from app.services.admin_service import OperationLogService
from app.services.admin_permission_service import require_permission
from app.services.chat_message_navigation_service import ChatMessageNavigationService
router = APIRouter()
@@ -96,7 +98,13 @@ def export_feedback(
@router.get("/admin/{feedback_id}")
def feedback_detail(feedback_id: int, db: Session = Depends(get_db), admin: Admin = Depends(get_current_admin)) -> dict:
def feedback_detail(
feedback_id: int,
messagePage: Annotated[int | None, Query(ge=1)] = None,
messagePageSize: Annotated[int, Query(ge=10, le=100)] = 20,
db: Session = Depends(get_db),
admin: Admin = Depends(get_current_admin),
) -> dict:
require_permission(admin, "feedback.detail")
row = db.execute(select(MessageFeedback, User, ChatMessage, ChatSession).join(User, User.id == MessageFeedback.user_id).join(ChatMessage, ChatMessage.id == MessageFeedback.message_id).join(ChatSession, ChatSession.id == MessageFeedback.session_id).where(MessageFeedback.id == feedback_id)).first()
if row is None:
@@ -105,10 +113,43 @@ def feedback_detail(feedback_id: int, db: Session = Depends(get_db), admin: Admi
if not feedback.is_read:
feedback.is_read = 1
feedback.read_by = admin.id
feedback.read_at = datetime.now(UTC).replace(tzinfo=None)
feedback.read_at = utc_now_naive()
db.commit()
messages = db.scalars(select(ChatMessage).where(ChatMessage.session_id == feedback.session_id, ChatMessage.id <= target.id).order_by(ChatMessage.id.asc()).limit(200)).all()
return api_success({**_summary(feedback, user, target), "sessionTitle": session.title, "messages": [{"id": m.id, "role": m.role, "content": m.content, "createdAt": m.created_at, "isTarget": m.id == target.id} for m in messages]})
message_query = select(ChatMessage).where(ChatMessage.session_id == feedback.session_id)
message_total = db.scalar(select(func.count()).select_from(message_query.subquery())) or 0
resolved_page = messagePage or ChatMessageNavigationService.page_for_message(
db,
session_id=feedback.session_id,
message=target,
page_size=messagePageSize,
)
messages = db.scalars(
message_query
.order_by(ChatMessage.created_at.asc(), ChatMessage.id.asc())
.offset((resolved_page - 1) * messagePageSize)
.limit(messagePageSize)
).all()
serialized_messages = [
{
"id": message.id,
"role": message.role,
"content": message.content,
"createdAt": message.created_at,
"isTarget": message.id == target.id,
}
for message in messages
]
return api_success({
**_summary(feedback, user, target),
"sessionTitle": session.title,
"messages": serialized_messages,
"messagesPage": page_result(
serialized_messages,
total=message_total,
page=resolved_page,
page_size=messagePageSize,
),
})
@router.delete("/admin/{feedback_id}")

View File

@@ -41,6 +41,9 @@ class ChatSession(Base, TimestampMixin):
class ChatMessage(Base):
__tablename__ = "sys_chat_message"
__table_args__ = (
Index("ix_chat_message_session_created_id", "session_id", "created_at", "id"),
)
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
session_id: Mapped[int] = mapped_column(ForeignKey("sys_chat_session.id"), index=True, nullable=False)

View File

@@ -0,0 +1,32 @@
from __future__ import annotations
from sqlalchemy import func, or_, select
from sqlalchemy.orm import Session
from app.models.chat import ChatMessage
class ChatMessageNavigationService:
"""Shared chronological paging rules for locating a message in a conversation."""
@staticmethod
def page_for_message(
db: Session,
*,
session_id: int,
message: ChatMessage,
page_size: int,
) -> int:
position = db.scalar(
select(func.count()).select_from(ChatMessage).where(
ChatMessage.session_id == session_id,
or_(
ChatMessage.created_at < message.created_at,
(
(ChatMessage.created_at == message.created_at)
& (ChatMessage.id <= message.id)
),
),
)
) or 1
return (position - 1) // page_size + 1

View File

@@ -99,6 +99,7 @@ def test_chat_detail_messages_are_paginated():
ChatMessage(id=index + 1, session_id=1, user_id=1, role="user", content=f"消息{index + 1}")
for index in range(25)
])
db.add(AiRequestLog(session_id=1, status="success", prompt="p" * 100_000, retrieved_chunks='[{"content":"large"}]'))
db.commit()
response = chat_detail(
@@ -116,6 +117,7 @@ def test_chat_detail_messages_are_paginated():
assert data["messagesPage"]["page"] == 2
assert len(data["messages"]) == 10
assert data["messages"][0]["content"] == "消息11"
assert data["aiLogs"] == []
def test_chat_detail_focus_message_opens_its_page_in_chronological_order():
@@ -164,7 +166,7 @@ def test_chat_messages_endpoint_returns_only_message_page():
db.add(AiRequestLog(session_id=1, status="success", prompt="p" * 5000, retrieved_chunks='[{"content":"large"}]'))
db.commit()
response = chat_messages(1, page=2, pageSize=10, db=db, current_admin=object())
response = chat_messages(1, keyword="", page=2, pageSize=10, db=db, current_admin=object())
data = response["data"]
assert data["total"] == 12
@@ -173,6 +175,32 @@ def test_chat_messages_endpoint_returns_only_message_page():
assert "aiLogs" not in data
def test_chat_messages_keyword_search_is_scoped_and_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=24)
other_session = ChatSession(id=2, user_id=1, title="其他会话", message_count=1)
db.add_all([user, session, other_session])
db.add_all([
ChatMessage(
id=index,
session_id=1,
user_id=1,
role="user",
content=f"{index}条 关键字" if index % 2 == 0 else f"{index}条普通内容",
)
for index in range(1, 25)
])
db.add(ChatMessage(id=100, session_id=2, user_id=1, role="user", content="其他会话关键字"))
db.commit()
response = chat_messages(1, keyword="关键字", page=2, pageSize=10, db=db, current_admin=object())
data = response["data"]
assert data["total"] == 12
assert [item["id"] for item in data["items"]] == [22, 24]
def test_retrieval_and_attention_lists_are_paginated():
with _database() as db:
for index in range(21):

View File

@@ -55,3 +55,47 @@ def test_feedback_export_workbook_is_formatted_and_formula_safe() -> None:
assert "AI生成内容请结合实际情况核对后使用。" in sheet["H2"].value
assert sheet["I2"].value == created_at + timedelta(hours=8)
assert sheet.tables["FeedbackRecords"].ref == "A1:J2"
def test_feedback_detail_locates_target_beyond_first_two_hundred_messages() -> None:
engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool)
Base.metadata.create_all(engine)
with Session(engine) as db:
user = User(id=1, phone="13800000000", name="测试用户")
admin = Admin(id=1, username="admin", password="hash", name="管理员", is_super_admin=1, must_change_password=0)
session = ChatSession(id=1, user_id=1, title="超长会话", message_count=250)
created_at = datetime(2026, 8, 29, 9, 0)
db.add_all([user, admin, session])
db.add_all([
ChatMessage(
id=index,
session_id=1,
user_id=1,
role="assistant" if index % 2 == 0 else "user",
content=f"消息{index}",
message_status="FINISHED",
created_at=created_at,
)
for index in range(1, 251)
])
db.commit()
feedback_id = create_feedback(
FeedbackCreate(messageId=230, content="这条回答有问题"),
db=db,
user=user,
)["data"]["id"]
detail = feedback_detail(
feedback_id,
messagePage=None,
messagePageSize=20,
db=db,
admin=admin,
)["data"]
assert detail["messagesPage"]["page"] == 12
assert detail["messagesPage"]["total"] == 250
assert [item["id"] for item in detail["messages"]] == list(range(221, 241))
target = next(item for item in detail["messages"] if item["isTarget"])
assert target["id"] == 230
assert target["createdAt"] == "2026-08-29T09:00:00.000Z"