diff --git a/internal/payload/incoming/messages.go b/internal/payload/incoming/messages.go index 1a7e8edb5..3cfa0aaef 100644 --- a/internal/payload/incoming/messages.go +++ b/internal/payload/incoming/messages.go @@ -22,4 +22,10 @@ type ( LastID uint64 `json:"lastID,string"` RepliesTo uint64 `json:"repliesTo,string"` } + + MessageThreads struct { + ChannelID uint64 `json:"channelId,string"` + FirstID uint64 `json:"firstID,string"` + LastID uint64 `json:"lastID,string"` + } ) diff --git a/internal/payload/incoming/payload.go b/internal/payload/incoming/payload.go index 1a4f2a5a9..94f5dd9f9 100644 --- a/internal/payload/incoming/payload.go +++ b/internal/payload/incoming/payload.go @@ -15,7 +15,8 @@ type Payload struct { *ChannelActivity `json:"channelActivity"` // Get channel message history - *Messages `json:"messages"` + *Messages `json:"messages"` + *MessageThreads `json:"messageThreads"` // Message actions *MessageCreate `json:"createMessage"` diff --git a/sam/repository/message.go b/sam/repository/message.go index b738faf93..055b0706a 100644 --- a/sam/repository/message.go +++ b/sam/repository/message.go @@ -15,6 +15,7 @@ type ( FindMessageByID(id uint64) (*types.Message, error) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) + FindThreads(filter *types.MessageFilter) (types.MessageSet, error) CreateMessage(mod *types.Message) (*types.Message, error) UpdateMessage(mod *types.Message) (*types.Message, error) DeleteMessageByID(ID uint64) error @@ -30,21 +31,37 @@ type ( const ( MESSAGES_MAX_LIMIT = 100 + sqlMessageColumns = "id, " + + "COALESCE(type,'') AS type, " + + "message, " + + "rel_user, " + + "rel_channel, " + + "reply_to, " + + "replies, " + + "created_at, " + + "updated_at, " + + "deleted_at" sqlMessageScope = "deleted_at IS NULL" - sqlMessagesSelect = `SELECT id, - COALESCE(type,'') AS type, - message, - rel_user, - rel_channel, - reply_to, - replies, - created_at, - updated_at, - deleted_at + sqlMessagesSelect = `SELECT ` + sqlMessageColumns + ` FROM messages WHERE ` + sqlMessageScope + sqlMessagesThreads = "WITH originals AS (" + + " SELECT id AS original_id " + + " FROM messages " + + " WHERE " + sqlMessageScope + + " AND rel_channel IN " + sqlChannelAccess + + " AND reply_to = 0 " + + " AND replies > 0 " + + " ORDER BY id DESC " + + " LIMIT ? " + + ")" + + " SELECT " + sqlMessageColumns + + " FROM messages, originals " + + " WHERE " + sqlMessageScope + + " AND original_id IN (id, reply_to)" + sqlMessageRepliesIncCount = `UPDATE messages SET replies = replies + 1 WHERE id = ? AND reply_to = 0` sqlMessageRepliesDecCount = `UPDATE messages SET replies = replies - 1 WHERE id = ? AND reply_to = 0` @@ -69,15 +86,13 @@ func (r *message) FindMessageByID(id uint64) (*types.Message, error) { } func (r *message) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) { + r.sanitizeFilter(filter) + params := make([]interface{}, 0) rval := make(types.MessageSet, 0) sql := sqlMessagesSelect - if filter == nil { - filter = &types.MessageFilter{} - } - if filter.Query != "" { sql += " AND message LIKE ?" params = append(params, filter.Query+"%") @@ -113,16 +128,43 @@ func (r *message) FindMessages(filter *types.MessageFilter) (types.MessageSet, e sql += " ORDER BY id DESC" - if filter.Limit == 0 || filter.Limit > MESSAGES_MAX_LIMIT { - filter.Limit = MESSAGES_MAX_LIMIT - } - sql += " LIMIT ? " params = append(params, filter.Limit) return rval, r.db().Select(&rval, sql, params...) } +func (r *message) FindThreads(filter *types.MessageFilter) (types.MessageSet, error) { + r.sanitizeFilter(filter) + + params := make([]interface{}, 0) + rval := make(types.MessageSet, 0) + + // for sqlChannelAccess + params = append(params, filter.CurrentUserID, types.ChannelTypePublic) + + // for sqlMessagesThreads + params = append(params, filter.Limit) + + sql := sqlMessagesThreads + if filter.ChannelID > 0 { + sql += " AND rel_channel = ? " + params = append(params, filter.ChannelID) + } + + return rval, r.db().Select(&rval, sql, params...) +} + +func (r *message) sanitizeFilter(filter *types.MessageFilter) { + if filter == nil { + filter = &types.MessageFilter{} + } + + if filter.Limit == 0 || filter.Limit > MESSAGES_MAX_LIMIT { + filter.Limit = MESSAGES_MAX_LIMIT + } +} + func (r *message) CreateMessage(mod *types.Message) (*types.Message, error) { mod.ID = factory.Sonyflake.NextID() mod.CreatedAt = time.Now() diff --git a/sam/repository/message_test.go b/sam/repository/message_test.go index 3b45979b4..5cdae5a44 100644 --- a/sam/repository/message_test.go +++ b/sam/repository/message_test.go @@ -95,6 +95,10 @@ func TestReplies(t *testing.T) { assert(t, err == nil, "CreateMessage error: %v", err) assert(t, rpl.ID > 0, "Reply did not get its ID") + // Let's increase this so that FindThreads + // can include it into results + msgRpo.IncReplyCount(msg.ID) + { mm, err = msgRpo.FindMessages(&types.MessageFilter{ RepliesTo: msg.ID, @@ -106,6 +110,17 @@ func TestReplies(t *testing.T) { assert(t, mm[0].ID == rpl.ID, "Reply ID does not match") } + { + mm, err = msgRpo.FindThreads(&types.MessageFilter{ + ChannelID: ch.ID, + }) + + assert(t, err == nil, "FindThreads error: %v", err) + assert(t, len(mm) == 2, "Failed to fetch messages in threads (2 messages), got: %d", len(mm)) + assert(t, mm[0].ID == msg.ID, "Original message ID does not match") + assert(t, mm[1].ID == rpl.ID, "Reply ID does not match") + } + { mm, err = msgRpo.FindMessages(&types.MessageFilter{ ChannelID: ch.ID, @@ -117,20 +132,20 @@ func TestReplies(t *testing.T) { } { + assert(t, msgRpo.IncReplyCount(msg.ID) == nil, "IncReplyCount should not return an error") assert(t, msgRpo.IncReplyCount(msg.ID) == nil, "IncReplyCount should not return an error") - assert(t, msgRpo.IncReplyCount(msg.ID) == nil, "IncReplyCount should not return an error") + // +1 that we have from before msg, err = msgRpo.FindMessageByID(msg.ID) assert(t, err == nil, "FindMessageByID error: %v", err) assert(t, msg.Replies == 3, "Reply counter check failed, expecting 3, got %v", msg.Replies) - assert(t, msgRpo.DecReplyCount(msg.ID) == nil, "DecReplyCount should not return an error") assert(t, msgRpo.DecReplyCount(msg.ID) == nil, "DecReplyCount should not return an error") msg, err = msgRpo.FindMessageByID(msg.ID) assert(t, err == nil, "FindMessageByID error: %v", err) - assert(t, msg.Replies == 1, "Reply counter check failed, expecting 1, got %v", msg.Replies) + assert(t, msg.Replies == 2, "Reply counter check failed, expecting 1, got %v", msg.Replies) } return nil diff --git a/sam/service/message.go b/sam/service/message.go index 354f37239..a43b91e9c 100644 --- a/sam/service/message.go +++ b/sam/service/message.go @@ -34,6 +34,7 @@ type ( With(ctx context.Context) MessageService Find(filter *types.MessageFilter) (types.MessageSet, error) + FindThreads(filter *types.MessageFilter) (types.MessageSet, error) Create(messages *types.Message) (*types.Message, error) Update(messages *types.Message) (*types.Message, error) @@ -93,6 +94,23 @@ func (svc *message) Find(filter *types.MessageFilter) (mm types.MessageSet, err return mm, svc.preloadAttachments(mm) } +func (svc *message) FindThreads(filter *types.MessageFilter) (mm types.MessageSet, err error) { + // @todo get user from context + filter.CurrentUserID = repository.Identity(svc.ctx) + + // @todo verify if current user can access & read from this channel + _ = filter.ChannelID + + mm, err = svc.message.FindThreads(filter) + if err != nil { + return nil, err + } + + svc.preloadUsers(mm) + + return mm, svc.preloadAttachments(mm) +} + func (svc *message) Create(mod *types.Message) (message *types.Message, err error) { // @todo get user from context var currentUserID uint64 = repository.Identity(svc.ctx) diff --git a/sam/websocket/session_incoming.go b/sam/websocket/session_incoming.go index 613e7788f..2981cce3d 100644 --- a/sam/websocket/session_incoming.go +++ b/sam/websocket/session_incoming.go @@ -23,6 +23,8 @@ func (s *Session) dispatch(raw []byte) error { return s.messageDelete(ctx, p.MessageDelete) case p.Messages != nil: return s.messageHistory(ctx, p.Messages) + case p.MessageThreads != nil: + return s.messageThreads(ctx, p.MessageThreads) // channel actions case p.ChannelJoin != nil: diff --git a/sam/websocket/session_incoming_message.go b/sam/websocket/session_incoming_message.go index f417b3cc4..7414b7da9 100644 --- a/sam/websocket/session_incoming_message.go +++ b/sam/websocket/session_incoming_message.go @@ -57,3 +57,28 @@ func (s *Session) messageHistory(ctx context.Context, p *incoming.Messages) erro return nil } + +func (s *Session) messageThreads(ctx context.Context, p *incoming.MessageThreads) error { + var ( + filter = &types.MessageFilter{ + ChannelID: p.ChannelID, + FirstID: p.FirstID, + LastID: p.LastID, + + // Max no. of messages we will return + Limit: 50, + } + ) + + messages, err := s.svc.msg.With(ctx).FindThreads(filter) + if err != nil { + return err + } + + err = s.sendReply(payload.Messages(ctx, messages)) + if err != nil { + return err + } + + return nil +}