3
0

Add FindThreads and expose it on websocket endpoint

This commit is contained in:
Denis Arh
2018-10-29 13:46:21 +01:00
parent 8c2d91122d
commit f8c0ad36d1
7 changed files with 131 additions and 22 deletions
+6
View File
@@ -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"`
}
)
+2 -1
View File
@@ -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
View File
@@ -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()
+18 -3
View File
@@ -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
+18
View File
@@ -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)
+2
View File
@@ -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:
+25
View File
@@ -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
}