Files
gochat/internal/repository/message_repo.go
T

248 lines
8.7 KiB
Go

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
}