From 8c488b83b0d4b56f4dce5fce68650fdd3a6a232d Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Mon, 19 Nov 2018 08:52:32 +0100 Subject: [PATCH] Count unread messages on mark-as-unread --- sam/repository/message.go | 8 ++++++++ sam/service/message.go | 14 ++++++++++++-- 2 files changed, 20 insertions(+), 2 deletions(-) diff --git a/sam/repository/message.go b/sam/repository/message.go index 3c268b2a6..fbe582472 100644 --- a/sam/repository/message.go +++ b/sam/repository/message.go @@ -17,6 +17,7 @@ type ( FindMessageByID(id uint64) (*types.Message, error) FindMessages(filter *types.MessageFilter) (types.MessageSet, error) FindThreads(filter *types.MessageFilter) (types.MessageSet, error) + CountFromMessageID(channelID, threadID, messageID uint64) (uint32, error) PrefillThreadParticipants(mm types.MessageSet) error CreateMessage(mod *types.Message) (*types.Message, error) UpdateMessage(mod *types.Message) (*types.Message, error) @@ -66,6 +67,8 @@ const ( sqlThreadParticipantsByMessageID = "SELECT DISTINCT reply_to, rel_user FROM messages WHERE reply_to IN (?)" + sqlCountFromMessageID = "SELECT COUNT(1) AS count FROM messages WHERE rel_channel = ? AND reply_to = ? AND id > ?" + 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` @@ -185,6 +188,11 @@ func (r *message) FindThreads(filter *types.MessageFilter) (types.MessageSet, er return rval, r.db().Select(&rval, sql, params...) } +func (r *message) CountFromMessageID(channelID, threadID, messageID uint64) (uint32, error) { + rval := struct{ Count uint32 }{} + return rval.Count, r.db().Get(&rval, sqlCountFromMessageID, channelID, threadID, messageID) +} + func (r *message) PrefillThreadParticipants(mm types.MessageSet) error { var rval = []struct { ReplyTo uint64 `db:"reply_to"` diff --git a/sam/service/message.go b/sam/service/message.go index c1f27214d..1bf8f434a 100644 --- a/sam/service/message.go +++ b/sam/service/message.go @@ -315,16 +315,26 @@ func (svc *message) MarkAsUnread(messageID uint64) error { return svc.db.Transaction(func() (err error) { // Broadcast queue var message *types.Message + var count uint32 message, err = svc.message.FindMessageByID(messageID) if err != nil { return err } + count, err = svc.message.CountFromMessageID(message.ChannelID, message.ReplyTo, message.ID) + if err != nil { + return + } + + // Inc counter so that we take + // this message into account + count++ + if message.ReplyTo > 0 { - return svc.unreads.Record(currentUserID, message.ChannelID, message.ReplyTo, messageID, 0) + return svc.unreads.Record(currentUserID, message.ChannelID, message.ReplyTo, messageID, count) } else { - return svc.unreads.Record(currentUserID, message.ChannelID, 0, messageID, 0) + return svc.unreads.Record(currentUserID, message.ChannelID, 0, messageID, count) } }) }