feat: improve chat feedback navigation
This commit is contained in:
@@ -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")
|
||||
@@ -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],
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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):
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user