Add FindThreads and expose it on websocket endpoint
This commit is contained in:
@@ -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"`
|
||||
}
|
||||
)
|
||||
|
||||
@@ -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"`
|
||||
|
||||
+60
-18
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user