232 lines
10 KiB
Python
232 lines
10 KiB
Python
from __future__ import annotations
|
||
|
||
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
|
||
from openpyxl import Workbook
|
||
from openpyxl.styles import Alignment, Font, PatternFill
|
||
from openpyxl.worksheet.table import Table, TableStyleInfo
|
||
from pydantic import BaseModel, Field
|
||
from sqlalchemy import func, select
|
||
from sqlalchemy.exc import IntegrityError
|
||
from sqlalchemy.orm import Session
|
||
|
||
from app.core.ai_content_label import ensure_ai_generated_notice
|
||
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, 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()
|
||
|
||
|
||
class FeedbackCreate(BaseModel):
|
||
messageId: int = Field(gt=0)
|
||
content: str = Field(min_length=1, max_length=200)
|
||
|
||
|
||
@router.post("")
|
||
def create_feedback(payload: FeedbackCreate, db: Session = Depends(get_db), user: User = Depends(get_current_user)) -> dict:
|
||
message = db.scalar(select(ChatMessage).where(ChatMessage.id == payload.messageId, ChatMessage.user_id == user.id))
|
||
if message is None or message.role != "assistant" or message.message_status != "FINISHED":
|
||
raise HTTPException(status_code=404, detail="反馈的回答不存在")
|
||
content = payload.content.strip()
|
||
if not content:
|
||
raise HTTPException(status_code=400, detail="请填写反馈内容")
|
||
existing = db.scalar(select(MessageFeedback).where(MessageFeedback.user_id == user.id, MessageFeedback.message_id == message.id))
|
||
if existing:
|
||
raise HTTPException(status_code=409, detail="这条回答已经反馈过了")
|
||
item = MessageFeedback(user_id=user.id, session_id=message.session_id, message_id=message.id, content=content)
|
||
db.add(item)
|
||
try:
|
||
db.commit()
|
||
except IntegrityError as exc:
|
||
db.rollback()
|
||
raise HTTPException(status_code=409, detail="这条回答已经反馈过了") from exc
|
||
db.refresh(item)
|
||
return api_success({"id": item.id})
|
||
|
||
|
||
@router.get("/admin/list")
|
||
def feedback_list(readStatus: str = Query(default="all", pattern="^(all|read|unread)$"), startDate: date | None = None, endDate: date | None = None, page: int = Query(default=1, ge=1), pageSize: int = Query(default=20, ge=10, le=100), db: Session = Depends(get_db), _admin: Admin = Depends(get_current_admin)) -> dict:
|
||
require_permission(_admin, "feedback.view")
|
||
query = _feedback_query(readStatus, startDate, endDate)
|
||
total = db.scalar(select(func.count()).select_from(query.order_by(None).subquery())) or 0
|
||
rows = db.execute(query.order_by(MessageFeedback.created_at.desc()).offset((page - 1) * pageSize).limit(pageSize)).all()
|
||
return api_success(page_result([_summary(row[0], row[1], row[2]) for row in rows], total=total, page=page, page_size=pageSize))
|
||
|
||
|
||
@router.get("/admin/export")
|
||
def export_feedback(
|
||
startDate: date,
|
||
endDate: date,
|
||
readStatus: str = Query(default="all", pattern="^(all|read|unread)$"),
|
||
db: Session = Depends(get_db),
|
||
admin: Admin = Depends(get_current_admin),
|
||
) -> StreamingResponse:
|
||
require_permission(admin, "feedback.export")
|
||
if endDate < startDate:
|
||
raise HTTPException(status_code=400, detail="结束日期不能早于开始日期")
|
||
if endDate - startDate > timedelta(days=366):
|
||
raise HTTPException(status_code=400, detail="单次导出时间范围不能超过 366 天")
|
||
rows = db.execute(_feedback_query(readStatus, startDate, endDate).order_by(MessageFeedback.created_at.desc()).limit(100001)).all()
|
||
if len(rows) > 100000:
|
||
raise HTTPException(status_code=400, detail="导出数据超过 10 万条,请缩小时间范围")
|
||
workbook = _feedback_workbook(rows)
|
||
stream = BytesIO()
|
||
workbook.save(stream)
|
||
stream.seek(0)
|
||
OperationLogService.write(db, admin_id=admin.id, module="feedback", action="export")
|
||
db.commit()
|
||
filename = f"feedback_{startDate:%Y%m%d}_{endDate:%Y%m%d}.xlsx"
|
||
return StreamingResponse(
|
||
stream,
|
||
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||
headers={"Content-Disposition": f'attachment; filename="{filename}"'},
|
||
)
|
||
|
||
|
||
@router.get("/admin/{feedback_id}")
|
||
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:
|
||
raise HTTPException(status_code=404, detail="反馈不存在")
|
||
feedback, user, target, session = row
|
||
if not feedback.is_read:
|
||
feedback.is_read = 1
|
||
feedback.read_by = admin.id
|
||
feedback.read_at = utc_now_naive()
|
||
db.commit()
|
||
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}")
|
||
def delete_feedback(feedback_id: int, db: Session = Depends(get_db), admin: Admin = Depends(get_current_admin)) -> dict:
|
||
require_permission(admin, "feedback.delete")
|
||
item = db.get(MessageFeedback, feedback_id)
|
||
if item is None:
|
||
raise HTTPException(status_code=404, detail="反馈不存在")
|
||
db.delete(item)
|
||
OperationLogService.write(db, admin_id=admin.id, module="feedback", action="delete", target_id=feedback_id)
|
||
db.commit()
|
||
return api_success()
|
||
|
||
|
||
def _summary(item: MessageFeedback, user: User, message: ChatMessage) -> dict:
|
||
return {"id": item.id, "userId": user.id, "userName": user.name, "userPhone": user.phone, "messageId": message.id, "messageContent": message.content, "content": item.content, "isRead": bool(item.is_read), "createdAt": item.created_at}
|
||
|
||
|
||
def _feedback_query(read_status: str, start_date: date | None, end_date: date | None):
|
||
query = (
|
||
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)
|
||
)
|
||
if read_status != "all":
|
||
query = query.where(MessageFeedback.is_read == (1 if read_status == "read" else 0))
|
||
if start_date is not None:
|
||
query = query.where(MessageFeedback.created_at >= business_day_boundary(start_date))
|
||
if end_date is not None:
|
||
query = query.where(MessageFeedback.created_at < business_day_boundary(end_date, end_exclusive=True))
|
||
return query
|
||
|
||
|
||
def _feedback_workbook(rows: list[tuple]) -> Workbook:
|
||
workbook = Workbook()
|
||
sheet = workbook.active
|
||
sheet.title = "反馈记录"
|
||
sheet.sheet_view.showGridLines = False
|
||
headers = ["序号", "状态", "用户姓名", "手机号", "反馈内容", "会话标题", "消息ID", "对应AI回答(AI生成)", "提交时间", "阅读时间"]
|
||
sheet.append(headers)
|
||
for index, (feedback, user, message, session) in enumerate(rows, start=1):
|
||
sheet.append([
|
||
index,
|
||
"已读" if feedback.is_read else "未读",
|
||
_excel_safe_text(user.name),
|
||
_excel_safe_text(user.phone),
|
||
_excel_safe_text(feedback.content),
|
||
_excel_safe_text(session.title),
|
||
message.id,
|
||
_excel_safe_text(ensure_ai_generated_notice(message.content)),
|
||
to_business_naive(feedback.created_at),
|
||
to_business_naive(feedback.read_at),
|
||
])
|
||
header_fill = PatternFill("solid", fgColor="1F6F5F")
|
||
for cell in sheet[1]:
|
||
cell.fill = header_fill
|
||
cell.font = Font(color="FFFFFF", bold=True)
|
||
cell.alignment = Alignment(horizontal="center", vertical="center")
|
||
sheet.freeze_panes = "A2"
|
||
sheet.auto_filter.ref = f"A1:J{max(1, sheet.max_row)}"
|
||
widths = [8, 10, 16, 16, 38, 24, 12, 70, 20, 20]
|
||
for index, width in enumerate(widths, start=1):
|
||
sheet.column_dimensions[chr(64 + index)].width = width
|
||
for row in sheet.iter_rows(min_row=2):
|
||
row[4].alignment = Alignment(vertical="top", wrap_text=True)
|
||
row[7].alignment = Alignment(vertical="top", wrap_text=True)
|
||
for cell in (row[8], row[9]):
|
||
cell.number_format = "yyyy-mm-dd hh:mm:ss"
|
||
if sheet.max_row > 1:
|
||
table = Table(displayName="FeedbackRecords", ref=f"A1:J{sheet.max_row}")
|
||
table.tableStyleInfo = TableStyleInfo(name="TableStyleMedium2", showRowStripes=True, showFirstColumn=False, showLastColumn=False)
|
||
sheet.add_table(table)
|
||
return workbook
|
||
|
||
|
||
def _excel_safe_text(value: object | None) -> str:
|
||
text_value = str(value or "")
|
||
return f"'{text_value}" if text_value.startswith(("=", "+", "-", "@")) else text_value
|