message_repository.go 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141
  1. package repository
  2. import (
  3. "errors"
  4. "time"
  5. "github.com/2930134478/AI-CS/backend/models"
  6. "gorm.io/gorm"
  7. )
  8. // MessageRepository 封装与消息相关的数据库操作。
  9. type MessageRepository struct {
  10. db *gorm.DB
  11. }
  12. // NewMessageRepository 创建消息仓库实例。
  13. func NewMessageRepository(db *gorm.DB) *MessageRepository {
  14. return &MessageRepository{db: db}
  15. }
  16. // Create 新建一条消息记录。
  17. func (r *MessageRepository) Create(message *models.Message) error {
  18. return r.db.Create(message).Error
  19. }
  20. // ListByConversationID 按时间顺序查询会话中的全部消息。
  21. func (r *MessageRepository) ListByConversationID(conversationID uint) ([]models.Message, error) {
  22. var messages []models.Message
  23. if err := r.db.Where("conversation_id = ?", conversationID).Order("created_at asc").Find(&messages).Error; err != nil {
  24. return nil, err
  25. }
  26. return messages, nil
  27. }
  28. // LatestByConversationID 查询会话中最新的一条消息。
  29. func (r *MessageRepository) LatestByConversationID(conversationID uint) (*models.Message, error) {
  30. var message models.Message
  31. if err := r.db.Where("conversation_id = ?", conversationID).
  32. Order("created_at desc").
  33. First(&message).Error; err != nil {
  34. if errors.Is(err, gorm.ErrRecordNotFound) {
  35. return nil, nil
  36. }
  37. return nil, err
  38. }
  39. return &message, nil
  40. }
  41. // CountUnreadBySender 统计指定发送方的未读消息数量。
  42. func (r *MessageRepository) CountUnreadBySender(conversationID uint, senderIsAgent bool) (int64, error) {
  43. var count int64
  44. if err := r.db.Model(&models.Message{}).
  45. Where("conversation_id = ? AND sender_is_agent = ? AND is_read = ?", conversationID, senderIsAgent, false).
  46. Count(&count).Error; err != nil {
  47. return 0, err
  48. }
  49. return count, nil
  50. }
  51. // FindConversationIDsByContent 根据关键字查询包含该内容的会话 ID。
  52. func (r *MessageRepository) FindConversationIDsByContent(keyword string) ([]uint, error) {
  53. var ids []uint
  54. if err := r.db.Model(&models.Message{}).
  55. Where("content LIKE ?", keyword).
  56. Pluck("conversation_id", &ids).Error; err != nil {
  57. return nil, err
  58. }
  59. return ids, nil
  60. }
  61. // MarkMessagesRead 将指定发送方的未读消息标记为已读,并返回受影响的消息 ID 及时间。
  62. func (r *MessageRepository) MarkMessagesRead(conversationID uint, senderIsAgent bool) ([]uint, int64, time.Time, error) {
  63. var messageIDs []uint
  64. if err := r.db.Model(&models.Message{}).
  65. Where("conversation_id = ? AND sender_is_agent = ? AND is_read = ?", conversationID, senderIsAgent, false).
  66. Pluck("id", &messageIDs).Error; err != nil {
  67. return nil, 0, time.Time{}, err
  68. }
  69. if len(messageIDs) == 0 {
  70. return []uint{}, 0, time.Time{}, nil
  71. }
  72. now := time.Now()
  73. if err := r.db.Model(&models.Message{}).
  74. Where("id IN ?", messageIDs).
  75. Updates(map[string]interface{}{
  76. "is_read": true,
  77. "read_at": now,
  78. }).Error; err != nil {
  79. return nil, 0, time.Time{}, err
  80. }
  81. remaining, err := r.CountUnreadBySender(conversationID, senderIsAgent)
  82. if err != nil {
  83. return nil, 0, time.Time{}, nil
  84. }
  85. return messageIDs, remaining, now, nil
  86. }
  87. // HasAgentJoinMessage 检查该对话中是否已经存在该客服的加入消息。
  88. // 用于避免重复创建"xxx加入了会话"的系统消息。
  89. func (r *MessageRepository) HasAgentJoinMessage(conversationID uint, agentID uint, agentName string) (bool, error) {
  90. var count int64
  91. joinMessageContent := agentName + "加入了会话"
  92. if err := r.db.Model(&models.Message{}).
  93. Where("conversation_id = ? AND sender_id = ? AND sender_is_agent = ? AND message_type = ? AND content = ?",
  94. conversationID, agentID, true, "system_message", joinMessageContent).
  95. Count(&count).Error; err != nil {
  96. return false, err
  97. }
  98. return count > 0, nil
  99. }
  100. // HasVisitorMessageInHumanMode 检查对话中是否有访客在人工模式下发送的消息。
  101. // 用于判断对话是否应该显示在客服列表中。
  102. // 只有当 ChatMode == "human" 且存在访客发送的消息时,才应该显示。
  103. func (r *MessageRepository) HasVisitorMessageInHumanMode(conversationID uint) (bool, error) {
  104. var count int64
  105. // 查询是否有访客发送的消息(sender_is_agent = false)
  106. // 注意:这里不检查 ChatMode,因为 ChatMode 在 Conversation 表中
  107. // 这个方法只检查消息是否存在,ChatMode 的检查在 Service 层
  108. if err := r.db.Model(&models.Message{}).
  109. Where("conversation_id = ? AND sender_is_agent = ?", conversationID, false).
  110. Count(&count).Error; err != nil {
  111. return false, err
  112. }
  113. return count > 0, nil
  114. }
  115. // HasAgentParticipated 检查指定客服是否在指定会话中发送过消息。
  116. // 用于判断该会话是否应该出现在该客服的"My chats"列表中。
  117. func (r *MessageRepository) HasAgentParticipated(conversationID uint, agentID uint) (bool, error) {
  118. var count int64
  119. // 查询是否有该客服发送的消息(sender_is_agent = true AND sender_id = agentID)
  120. // 注意:系统消息(message_type = 'system_message')也应该算作参与
  121. // 所以不限制 message_type,包括所有类型的消息
  122. if err := r.db.Model(&models.Message{}).
  123. Where("conversation_id = ? AND sender_is_agent = ? AND sender_id = ?", conversationID, true, agentID).
  124. Count(&count).Error; err != nil {
  125. return false, err
  126. }
  127. return count > 0, nil
  128. }