package repository import ( "context" "gorm.io/gorm" "github.com/gochat/gochat/internal/model" "github.com/gochat/gochat/internal/search" ) // MessageRepo implements GORM repository for Message. // Reference: Chatwoot app/models/message.rb type MessageRepo struct { db *gorm.DB } // DB returns the underlying gorm.DB for advanced query building. func (r *MessageRepo) DB() *gorm.DB { return r.db } // NewMessageRepo creates a new Message repository. func NewMessageRepo(db *gorm.DB) *MessageRepo { return &MessageRepo{db: db} } // FindByID retrieves a message by primary key. func (r *MessageRepo) FindByID(ctx context.Context, id uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx).First(&message, id).Error if err != nil { return nil, err } return &message, nil } // FindByConversation retrieves all messages for a conversation. func (r *MessageRepo) FindByConversation(ctx context.Context, conversationID uint, offset, limit int) ([]model.Message, int64, error) { var messages []model.Message var total int64 countDB := r.db.WithContext(ctx).Model(&model.Message{}).Where("conversation_id = ?", conversationID) if err := countDB.Count(&total).Error; err != nil { return nil, 0, err } err := r.db.WithContext(ctx).Where("conversation_id = ?", conversationID). Offset(offset).Limit(limit).Order("id ASC"). Find(&messages).Error return messages, total, err } // FindByConversationFinder mirrors Chatwoot's MessageFinder before/after windows. func (r *MessageRepo) FindByConversationFinder(ctx context.Context, conversationID uint, after, before uint, filterInternal bool) ([]model.Message, int64, error) { query := r.db.WithContext(ctx).Model(&model.Message{}).Where("conversation_id = ?", conversationID) if filterInternal { query = query.Where("NOT (private = ? OR message_type = ?)", true, model.MessageTypeActivity) } var total int64 if err := query.Count(&total).Error; err != nil { return nil, 0, err } var messages []model.Message switch { case after != 0 && before != 0: err := query.Where("id >= ? AND id < ?", after, before).Order("created_at ASC, id ASC").Limit(1000).Find(&messages).Error return messages, total, err case before != 0: err := query.Where("id < ?", before).Order("created_at DESC, id DESC").Limit(20).Find(&messages).Error if err != nil { return nil, 0, err } reverseMessages(messages) return messages, total, nil case after != 0: err := query.Where("id > ?", after).Order("created_at ASC, id ASC").Limit(100).Find(&messages).Error return messages, total, err default: err := query.Order("created_at DESC, id DESC").Limit(20).Find(&messages).Error if err != nil { return nil, 0, err } reverseMessages(messages) return messages, total, nil } } func reverseMessages(messages []model.Message) { for i, j := 0, len(messages)-1; i < j; i, j = i+1, j-1 { messages[i], messages[j] = messages[j], messages[i] } } // FindByAccountAndID retrieves a message scoped to an account. func (r *MessageRepo) FindByAccountAndID(ctx context.Context, accountID, id uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx).Where("account_id = ? AND id = ?", accountID, id).First(&message).Error if err != nil { return nil, err } return &message, nil } // FindByConversationAndID retrieves a message scoped to a conversation. func (r *MessageRepo) FindByConversationAndID(ctx context.Context, conversationID, id uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx).Where("conversation_id = ? AND id = ?", conversationID, id).First(&message).Error if err != nil { return nil, err } return &message, nil } // Search performs a full-text search on messages for an account. func (r *MessageRepo) Search(ctx context.Context, accountID uint, query string, offset, limit int, searchMode search.SearchMode) ([]model.Message, int64, error) { var messages []model.Message var total int64 searchQuery := r.db.WithContext(ctx).Model(&model.Message{}).Where("account_id = ?", accountID) switch searchMode { case search.SearchModeILike: searchQuery = searchQuery.Where("content LIKE ?", "%"+query+"%") case search.SearchModeTrigram: searchQuery = searchQuery.Where("content LIKE ?", "%"+query+"%") default: searchQuery = searchQuery.Where("content LIKE ?", "%"+query+"%") } if err := searchQuery.Count(&total).Error; err != nil { return nil, 0, err } err := searchQuery.Offset(offset).Limit(limit).Order("id DESC").Find(&messages).Error return messages, total, err } // Create persists a new message. func (r *MessageRepo) Create(ctx context.Context, message *model.Message) error { return r.db.WithContext(ctx).Create(message).Error } // Update persists changes to an existing message. func (r *MessageRepo) Update(ctx context.Context, message *model.Message) error { return r.db.WithContext(ctx).Save(message).Error } // Delete soft-deletes a message. func (r *MessageRepo) Delete(ctx context.Context, id uint) error { return r.db.WithContext(ctx).Delete(&model.Message{}, id).Error } // CountByConversation returns the message count for a conversation. func (r *MessageRepo) CountByConversation(ctx context.Context, conversationID uint) (int64, error) { var count int64 err := r.db.WithContext(ctx).Model(&model.Message{}).Where("conversation_id = ?", conversationID).Count(&count).Error return count, err } // FindLastByConversation retrieves the most recent message in a conversation. func (r *MessageRepo) FindLastByConversation(ctx context.Context, conversationID uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx).Where("conversation_id = ?", conversationID). Order("id DESC").First(&message).Error if err != nil { return nil, err } return &message, nil } // FindByAccount retrieves all messages for an account with pagination. func (r *MessageRepo) FindByAccount(ctx context.Context, accountID uint, offset, limit int) ([]model.Message, int64, error) { var messages []model.Message var total int64 countDB := r.db.WithContext(ctx).Model(&model.Message{}).Where("account_id = ?", accountID) if err := countDB.Count(&total).Error; err != nil { return nil, 0, err } err := r.db.WithContext(ctx).Where("account_id = ?", accountID). Offset(offset).Limit(limit).Order("id DESC"). Find(&messages).Error return messages, total, err } // UpdateStatus updates the delivery status of a message. func (r *MessageRepo) UpdateStatus(ctx context.Context, id uint, status string) error { return r.db.WithContext(ctx).Model(&model.Message{}).Where("id = ?", id). Update("status", status).Error } // FindByStatus retrieves messages by status for a conversation in an account. func (r *MessageRepo) FindByStatus(ctx context.Context, accountID, conversationID uint, status string, offset, limit int) ([]model.Message, int64, error) { var messages []model.Message var total int64 countDB := r.db.WithContext(ctx).Model(&model.Message{}). Where("account_id = ? AND conversation_id = ? AND status = ?", accountID, conversationID, status) if err := countDB.Count(&total).Error; err != nil { return nil, 0, err } err := r.db.WithContext(ctx). Where("account_id = ? AND conversation_id = ? AND status = ?", accountID, conversationID, status). Offset(offset).Limit(limit).Order("id ASC"). Find(&messages).Error return messages, total, err } // FindLastMessage retrieves the most recent message in a conversation (alias for FindLastByConversation). func (r *MessageRepo) FindLastMessage(ctx context.Context, conversationID uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx).Where("conversation_id = ?", conversationID). Order("created_at DESC").First(&message).Error if err != nil { return nil, err } return &message, nil } // FindFirstAgentMessage retrieves the first outgoing/agent message in a conversation. func (r *MessageRepo) FindFirstAgentMessage(ctx context.Context, conversationID uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx). Where("conversation_id = ? AND message_type = ?", conversationID, model.MessageTypeOutgoing). Order("id ASC").First(&message).Error if err != nil { return nil, err } return &message, nil } // FindLastIncomingByConversation retrieves the last incoming message in a conversation. // Used by MarkUnread to set agent_last_seen_at to last_incoming.CreatedAt - 1s. func (r *MessageRepo) FindLastIncomingByConversation(ctx context.Context, conversationID uint) (*model.Message, error) { var message model.Message err := r.db.WithContext(ctx). Where("conversation_id = ? AND message_type = ?", conversationID, model.MessageTypeIncoming). Order("id DESC").First(&message).Error if err != nil { return nil, err } return &message, nil }