| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596 |
- package repository
- import (
- "errors"
- "time"
- "github.com/2930134478/AI-CS/backend/models"
- "gorm.io/gorm"
- )
- // MessageRepository 封装与消息相关的数据库操作。
- type MessageRepository struct {
- db *gorm.DB
- }
- // NewMessageRepository 创建消息仓库实例。
- func NewMessageRepository(db *gorm.DB) *MessageRepository {
- return &MessageRepository{db: db}
- }
- // Create 新建一条消息记录。
- func (r *MessageRepository) Create(message *models.Message) error {
- return r.db.Create(message).Error
- }
- // ListByConversationID 按时间顺序查询会话中的全部消息。
- func (r *MessageRepository) ListByConversationID(conversationID uint) ([]models.Message, error) {
- var messages []models.Message
- if err := r.db.Where("conversation_id = ?", conversationID).Order("created_at asc").Find(&messages).Error; err != nil {
- return nil, err
- }
- return messages, nil
- }
- // LatestByConversationID 查询会话中最新的一条消息。
- func (r *MessageRepository) LatestByConversationID(conversationID uint) (*models.Message, error) {
- var message models.Message
- if err := r.db.Where("conversation_id = ?", conversationID).
- Order("created_at desc").
- First(&message).Error; err != nil {
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, nil
- }
- return nil, err
- }
- return &message, nil
- }
- // CountUnreadBySender 统计指定发送方的未读消息数量。
- func (r *MessageRepository) CountUnreadBySender(conversationID uint, senderIsAgent bool) (int64, error) {
- var count int64
- if err := r.db.Model(&models.Message{}).
- Where("conversation_id = ? AND sender_is_agent = ? AND is_read = ?", conversationID, senderIsAgent, false).
- Count(&count).Error; err != nil {
- return 0, err
- }
- return count, nil
- }
- // FindConversationIDsByContent 根据关键字查询包含该内容的会话 ID。
- func (r *MessageRepository) FindConversationIDsByContent(keyword string) ([]uint, error) {
- var ids []uint
- if err := r.db.Model(&models.Message{}).
- Where("content LIKE ?", keyword).
- Pluck("conversation_id", &ids).Error; err != nil {
- return nil, err
- }
- return ids, nil
- }
- // MarkMessagesRead 将指定发送方的未读消息标记为已读,并返回受影响的消息 ID 及时间。
- func (r *MessageRepository) MarkMessagesRead(conversationID uint, senderIsAgent bool) ([]uint, int64, time.Time, error) {
- var messageIDs []uint
- if err := r.db.Model(&models.Message{}).
- Where("conversation_id = ? AND sender_is_agent = ? AND is_read = ?", conversationID, senderIsAgent, false).
- Pluck("id", &messageIDs).Error; err != nil {
- return nil, 0, time.Time{}, err
- }
- if len(messageIDs) == 0 {
- return []uint{}, 0, time.Time{}, nil
- }
- now := time.Now()
- if err := r.db.Model(&models.Message{}).
- Where("id IN ?", messageIDs).
- Updates(map[string]interface{}{
- "is_read": true,
- "read_at": now,
- }).Error; err != nil {
- return nil, 0, time.Time{}, err
- }
- remaining, err := r.CountUnreadBySender(conversationID, senderIsAgent)
- if err != nil {
- return nil, 0, time.Time{}, err
- }
- return messageIDs, remaining, now, nil
- }
|