message_controller.go 13 KB

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