message_controller.go 16 KB

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