message_controller.go 15 KB

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