message_controller.go 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370
  1. package controller
  2. import (
  3. "bytes"
  4. "io"
  5. "log"
  6. "net/http"
  7. "path/filepath"
  8. "strconv"
  9. "strings"
  10. "github.com/2930134478/AI-CS/backend/infra"
  11. "github.com/2930134478/AI-CS/backend/service"
  12. "github.com/gin-gonic/gin"
  13. )
  14. // MessageController 负责处理消息相关的 HTTP 请求。
  15. type MessageController struct {
  16. messageService *service.MessageService
  17. conversationService *service.ConversationService
  18. storageService infra.StorageService
  19. }
  20. // NewMessageController 创建 MessageController 实例。
  21. func NewMessageController(messageService *service.MessageService, conversationService *service.ConversationService, storageService infra.StorageService) *MessageController {
  22. return &MessageController{
  23. messageService: messageService,
  24. conversationService: conversationService,
  25. storageService: storageService,
  26. }
  27. }
  28. type createMessageRequest struct {
  29. ConversationID uint `json:"conversation_id"`
  30. Content string `json:"content"`
  31. SenderIsAgent bool `json:"sender_is_agent"`
  32. SenderID uint `json:"sender_id"`
  33. // 文件相关字段(可选)
  34. FileURL *string `json:"file_url"`
  35. FileType *string `json:"file_type"`
  36. FileName *string `json:"file_name"`
  37. FileSize *int64 `json:"file_size"`
  38. MimeType *string `json:"mime_type"`
  39. }
  40. // CreateMessage 处理发送消息的请求。
  41. func (mc *MessageController) CreateMessage(c *gin.Context) {
  42. var req createMessageRequest
  43. if err := c.ShouldBindJSON(&req); err != nil || req.ConversationID == 0 {
  44. c.JSON(http.StatusBadRequest, gin.H{"error": "请求参数错误"})
  45. return
  46. }
  47. // 验证:必须有内容或文件
  48. if req.Content == "" && req.FileURL == nil {
  49. c.JSON(http.StatusBadRequest, gin.H{"error": "消息内容或文件不能同时为空"})
  50. return
  51. }
  52. _, err := mc.messageService.CreateMessage(service.CreateMessageInput{
  53. ConversationID: req.ConversationID,
  54. Content: req.Content,
  55. SenderID: req.SenderID,
  56. SenderIsAgent: req.SenderIsAgent,
  57. FileURL: req.FileURL,
  58. FileType: req.FileType,
  59. FileName: req.FileName,
  60. FileSize: req.FileSize,
  61. MimeType: req.MimeType,
  62. })
  63. if err != nil {
  64. log.Printf("❌ 创建消息失败: 对话ID=%d, 错误=%v", req.ConversationID, err)
  65. switch err {
  66. case service.ErrConversationClosed:
  67. c.JSON(http.StatusBadRequest, gin.H{"error": "会话已关闭"})
  68. case service.ErrConversationNotFound:
  69. c.JSON(http.StatusBadRequest, gin.H{"error": "会话不存在"})
  70. default:
  71. c.JSON(http.StatusInternalServerError, gin.H{"error": "创建消息失败"})
  72. }
  73. return
  74. }
  75. c.JSON(http.StatusOK, gin.H{"message": "创建消息成功"})
  76. }
  77. // ListMessages 返回指定会话的消息列表。
  78. // 查询参数:
  79. // - conversation_id: 会话ID(必需)
  80. // - include_ai_messages: 是否包含 AI 消息(可选,默认 false)
  81. func (mc *MessageController) ListMessages(c *gin.Context) {
  82. conversationIDStr := c.Query("conversation_id")
  83. if conversationIDStr == "" {
  84. c.JSON(http.StatusBadRequest, gin.H{"error": "会话ID不能为空"})
  85. return
  86. }
  87. conversationID, err := strconv.ParseUint(conversationIDStr, 10, 64)
  88. if err != nil || conversationID == 0 {
  89. c.JSON(http.StatusBadRequest, gin.H{"error": "会话ID不合法"})
  90. return
  91. }
  92. // 解析 include_ai_messages 参数(默认 false)
  93. includeAIMessages := c.DefaultQuery("include_ai_messages", "false") == "true"
  94. messages, err := mc.messageService.ListMessages(uint(conversationID), includeAIMessages)
  95. if err != nil {
  96. c.JSON(http.StatusInternalServerError, gin.H{"error": "查询消息失败"})
  97. return
  98. }
  99. c.JSON(http.StatusOK, messages)
  100. }
  101. type markMessagesReadRequest struct {
  102. ConversationID uint `json:"conversation_id"`
  103. ReaderIsAgent bool `json:"reader_is_agent"`
  104. }
  105. // MarkMessagesRead 将指定会话的消息标记为已读。
  106. func (mc *MessageController) MarkMessagesRead(c *gin.Context) {
  107. var req markMessagesReadRequest
  108. if err := c.ShouldBindJSON(&req); err != nil || req.ConversationID == 0 {
  109. c.JSON(http.StatusBadRequest, gin.H{"error": "请求参数错误"})
  110. return
  111. }
  112. result, err := mc.messageService.MarkMessagesRead(req.ConversationID, req.ReaderIsAgent)
  113. if err != nil {
  114. c.JSON(http.StatusInternalServerError, gin.H{"error": "更新消息状态失败"})
  115. return
  116. }
  117. c.JSON(http.StatusOK, gin.H{
  118. "updated": len(result.MessageIDs),
  119. "message_ids": result.MessageIDs,
  120. "conversation_id": result.ConversationID,
  121. "unread_count": result.UnreadCount,
  122. "read_at": formatTimeValue(result.ReadAt),
  123. })
  124. }
  125. // UploadFile 处理文件上传请求。
  126. // 请求格式:multipart/form-data
  127. // - file: 文件内容(必需)
  128. // - conversation_id: 对话ID(可选,用于组织目录)
  129. // 认证方式:
  130. // - 方式1:提供 X-User-Id 请求头(客服上传)
  131. // - 方式2:提供 conversation_id 参数(访客上传,会验证对话是否存在且未关闭)
  132. func (mc *MessageController) UploadFile(c *gin.Context) {
  133. // ⚠️ 认证检查:必须满足以下条件之一
  134. // 1. 提供 X-User-Id 请求头(客服)
  135. // 2. 提供 conversation_id 参数(访客)
  136. userID := getUserIDFromHeader(c)
  137. conversationIDStr := c.PostForm("conversation_id")
  138. // 如果既没有用户ID,也没有对话ID,拒绝访问
  139. if userID == 0 && conversationIDStr == "" {
  140. c.JSON(http.StatusUnauthorized, gin.H{"error": "未授权访问,请提供 X-User-Id 请求头或 conversation_id 参数"})
  141. return
  142. }
  143. // 如果是访客上传(没有用户ID,但有对话ID),验证对话是否存在且未关闭
  144. if userID == 0 && conversationIDStr != "" {
  145. convID, err := strconv.ParseUint(conversationIDStr, 10, 64)
  146. if err != nil || convID == 0 {
  147. c.JSON(http.StatusBadRequest, gin.H{"error": "对话ID不合法"})
  148. return
  149. }
  150. // 验证对话是否存在且未关闭
  151. conv, err := mc.conversationService.GetConversationDetail(uint(convID), 0)
  152. if err != nil {
  153. c.JSON(http.StatusForbidden, gin.H{"error": "对话不存在或已关闭"})
  154. return
  155. }
  156. if conv.Status == "closed" {
  157. c.JSON(http.StatusForbidden, gin.H{"error": "对话已关闭"})
  158. return
  159. }
  160. }
  161. // 解析文件
  162. file, err := c.FormFile("file")
  163. if err != nil {
  164. c.JSON(http.StatusBadRequest, gin.H{"error": "文件不能为空"})
  165. return
  166. }
  167. // 验证文件大小(10MB)
  168. const maxFileSize = 10 * 1024 * 1024 // 10MB
  169. if file.Size > maxFileSize {
  170. c.JSON(http.StatusBadRequest, gin.H{"error": "文件大小超过限制(最大10MB)"})
  171. return
  172. }
  173. // ⚠️ 加强:验证文件类型(扩展名)
  174. ext := strings.ToLower(filepath.Ext(file.Filename))
  175. allowedExts := map[string]bool{
  176. ".jpg": true,
  177. ".jpeg": true,
  178. ".png": true,
  179. ".gif": true,
  180. ".webp": true,
  181. ".pdf": true,
  182. ".doc": true,
  183. ".docx": true,
  184. ".txt": true,
  185. }
  186. if !allowedExts[ext] {
  187. c.JSON(http.StatusBadRequest, gin.H{"error": "不支持的文件类型"})
  188. return
  189. }
  190. // ⚠️ 加强:验证 MIME 类型(防止伪造扩展名)
  191. mimeType := file.Header.Get("Content-Type")
  192. allowedMimeTypes := map[string]bool{
  193. "image/jpeg": true,
  194. "image/jpg": true,
  195. "image/png": true,
  196. "image/gif": true,
  197. "image/webp": true,
  198. "application/pdf": true,
  199. "application/msword": true,
  200. "application/vnd.openxmlformats-officedocument.wordprocessingml.document": true, // .docx
  201. "text/plain": true,
  202. }
  203. if !allowedMimeTypes[mimeType] {
  204. c.JSON(http.StatusBadRequest, gin.H{"error": "不支持的文件 MIME 类型: " + mimeType})
  205. return
  206. }
  207. // ⚠️ 加强:清理文件名,防止路径遍历攻击
  208. safeFilename := filepath.Base(file.Filename)
  209. safeFilename = strings.ReplaceAll(safeFilename, "..", "")
  210. safeFilename = strings.ReplaceAll(safeFilename, "/", "")
  211. safeFilename = strings.ReplaceAll(safeFilename, "\\", "")
  212. // 移除所有非字母数字、点、下划线、连字符的字符
  213. var cleaned strings.Builder
  214. for _, r := range safeFilename {
  215. if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9') || r == '.' || r == '_' || r == '-' {
  216. cleaned.WriteRune(r)
  217. }
  218. }
  219. safeFilename = cleaned.String()
  220. // 限制文件名长度
  221. if len(safeFilename) > 100 {
  222. // 保留扩展名
  223. ext := filepath.Ext(safeFilename)
  224. nameWithoutExt := strings.TrimSuffix(safeFilename, ext)
  225. if len(nameWithoutExt) > 100-len(ext) {
  226. safeFilename = nameWithoutExt[:100-len(ext)] + ext
  227. }
  228. }
  229. // ⚠️ 加强:验证文件内容(magic number 检查,防止伪造扩展名)
  230. fileContent, err := file.Open()
  231. if err != nil {
  232. c.JSON(http.StatusBadRequest, gin.H{"error": "无法读取文件"})
  233. return
  234. }
  235. defer fileContent.Close()
  236. // 读取文件前几个字节(magic number)
  237. magicBytes := make([]byte, 12)
  238. n, err := fileContent.Read(magicBytes)
  239. if err != nil && err != io.EOF {
  240. c.JSON(http.StatusBadRequest, gin.H{"error": "无法读取文件内容"})
  241. return
  242. }
  243. // 验证文件内容是否匹配扩展名
  244. if !isValidFileContent(ext, magicBytes[:n]) {
  245. c.JSON(http.StatusBadRequest, gin.H{"error": "文件内容与扩展名不匹配,可能是伪造的文件类型"})
  246. return
  247. }
  248. // 重置文件指针,以便后续保存
  249. if _, err := fileContent.Seek(0, io.SeekStart); err != nil {
  250. c.JSON(http.StatusInternalServerError, gin.H{"error": "无法重置文件指针"})
  251. return
  252. }
  253. // 获取对话ID(如果之前已经解析过,直接使用;否则从表单获取)
  254. var conversationID uint
  255. if conversationIDStr != "" {
  256. if id, err := strconv.ParseUint(conversationIDStr, 10, 64); err == nil {
  257. conversationID = uint(id)
  258. }
  259. }
  260. // 保存文件(使用清理后的文件名,fileContent 已经在上面打开并验证过)
  261. fileURL, err := mc.storageService.SaveMessageFile(conversationID, fileContent, safeFilename)
  262. if err != nil {
  263. log.Printf("❌ 保存文件失败: %v", err)
  264. c.JSON(http.StatusInternalServerError, gin.H{"error": "保存文件失败"})
  265. return
  266. }
  267. // 判断文件类型
  268. fileType := "document"
  269. if strings.HasPrefix(mimeType, "image/") {
  270. fileType = "image"
  271. }
  272. // 返回文件信息(使用清理后的文件名)
  273. c.JSON(http.StatusOK, gin.H{
  274. "success": true,
  275. "data": gin.H{
  276. "file_url": fileURL,
  277. "file_type": fileType,
  278. "file_name": safeFilename,
  279. "file_size": file.Size,
  280. "mime_type": mimeType,
  281. },
  282. })
  283. }
  284. // isValidFileContent 验证文件内容是否与扩展名匹配(通过 magic number 检查)
  285. func isValidFileContent(ext string, magicBytes []byte) bool {
  286. if len(magicBytes) < 4 {
  287. return false
  288. }
  289. ext = strings.ToLower(ext)
  290. // 检查各种文件类型的 magic number
  291. switch ext {
  292. case ".jpg", ".jpeg":
  293. // JPEG: FF D8 FF
  294. return len(magicBytes) >= 3 && magicBytes[0] == 0xFF && magicBytes[1] == 0xD8 && magicBytes[2] == 0xFF
  295. case ".png":
  296. // PNG: 89 50 4E 47
  297. return len(magicBytes) >= 4 && magicBytes[0] == 0x89 && magicBytes[1] == 0x50 && magicBytes[2] == 0x4E && magicBytes[3] == 0x47
  298. case ".gif":
  299. // GIF: 47 49 46 38 (GIF8)
  300. return len(magicBytes) >= 4 && magicBytes[0] == 0x47 && magicBytes[1] == 0x49 && magicBytes[2] == 0x46 && magicBytes[3] == 0x38
  301. case ".webp":
  302. // WebP: RIFF ... WEBP
  303. if len(magicBytes) >= 12 {
  304. return bytes.Equal(magicBytes[0:4], []byte("RIFF")) && bytes.Equal(magicBytes[8:12], []byte("WEBP"))
  305. }
  306. return false
  307. case ".pdf":
  308. // PDF: 25 50 44 46 (%PDF)
  309. return len(magicBytes) >= 4 && magicBytes[0] == 0x25 && magicBytes[1] == 0x50 && magicBytes[2] == 0x44 && magicBytes[3] == 0x46
  310. case ".txt":
  311. // 文本文件:检查是否为可打印字符(ASCII 32-126)或 UTF-8 BOM
  312. // UTF-8 BOM: EF BB BF
  313. if len(magicBytes) >= 3 && magicBytes[0] == 0xEF && magicBytes[1] == 0xBB && magicBytes[2] == 0xBF {
  314. return true
  315. }
  316. // 检查前几个字节是否都是可打印字符
  317. for i := 0; i < len(magicBytes) && i < 10; i++ {
  318. if magicBytes[i] < 0x20 && magicBytes[i] != 0x09 && magicBytes[i] != 0x0A && magicBytes[i] != 0x0D {
  319. // 不是可打印字符、制表符、换行符或回车符
  320. return false
  321. }
  322. }
  323. return true
  324. case ".doc":
  325. // DOC (OLE2): D0 CF 11 E0 A1 B1 1A E1
  326. return len(magicBytes) >= 8 && magicBytes[0] == 0xD0 && magicBytes[1] == 0xCF && magicBytes[2] == 0x11 && magicBytes[3] == 0xE0
  327. case ".docx":
  328. // DOCX (ZIP): 50 4B 03 04 (PK..)
  329. return len(magicBytes) >= 4 && magicBytes[0] == 0x50 && magicBytes[1] == 0x4B && magicBytes[2] == 0x03 && magicBytes[3] == 0x04
  330. default:
  331. return false
  332. }
  333. }