from __future__ import annotations from datetime import UTC, date, datetime, time, timedelta from io import BytesIO 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.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.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 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, 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 = datetime.now(UTC).replace(tzinfo=None) 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]}) @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 >= datetime.combine(start_date, time.min)) if end_date is not None: query = query.where(MessageFeedback.created_at < datetime.combine(end_date + timedelta(days=1), time.min)) return query def _feedback_workbook(rows: list[tuple]) -> Workbook: workbook = Workbook() sheet = workbook.active sheet.title = "反馈记录" sheet.sheet_view.showGridLines = False headers = ["序号", "状态", "用户姓名", "手机号", "反馈内容", "会话标题", "消息ID", "对应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(message.content), feedback.created_at, 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