diff --git a/messaging/internal/repository/unread.go b/messaging/internal/repository/unread.go index 2bd434a2b..a7534ef1b 100644 --- a/messaging/internal/repository/unread.go +++ b/messaging/internal/repository/unread.go @@ -33,7 +33,7 @@ const ( // Fetching channel members of all channels a specific user has access to sqlUnreadSelect = `SELECT rel_channel, rel_reply_to, rel_user, count, rel_last_message FROM messaging_unread - WHERE count > 0 ` + WHERE count > 0 && rel_last_message > 0 ` // Fetching channel members of all channels a specific user has access to sqlThreadUnreadSelect = `SELECT rel_channel, sum(count) as count @@ -81,6 +81,12 @@ func (r *unread) Find(filter *types.UnreadFilter) (uu types.UnreadSet, err error params = append(params, filter.UserID) } + if filter.ChannelID > 0 { + // scope: only channel we have access to + sql += ` AND rel_channel = ?` + params = append(params, filter.ChannelID) + } + if len(filter.ThreadIDs) > 0 { sql += ` AND rel_reply_to IN (?)` params = append(params, filter.ThreadIDs) diff --git a/messaging/internal/service/message.go b/messaging/internal/service/message.go index bc868c359..f54bfc60f 100644 --- a/messaging/internal/service/message.go +++ b/messaging/internal/service/message.go @@ -56,7 +56,7 @@ type ( React(messageID uint64, reaction string) error RemoveReaction(messageID uint64, reaction string) error - MarkAsRead(channelID, threadID, lastReadMessageID uint64) (uint64, uint32, error) + MarkAsRead(channelID, threadID, lastReadMessageID uint64) (uint64, uint32, uint32, error) Pin(messageID uint64) error RemovePin(messageID uint64) error @@ -431,10 +431,11 @@ func (svc message) Delete(messageID uint64) error { // MarkAsRead marks channel/thread as read // // If lastReadMessageID is set, it uses that message as last read message -func (svc message) MarkAsRead(channelID, threadID, lastReadMessageID uint64) (uint64, uint32, error) { +func (svc message) MarkAsRead(channelID, threadID, lastReadMessageID uint64) (uint64, uint32, uint32, error) { var ( currentUserID uint64 = repository.Identity(svc.ctx) count uint32 + tcount uint32 err error ) @@ -458,6 +459,14 @@ func (svc message) MarkAsRead(channelID, threadID, lastReadMessageID uint64) (ui } else if !thread.IsValid() { return errors.New("invalid thread") } + } else { + // This is request for channel, + // count all thread unreads + var uu types.UnreadSet + uu, err = svc.unread.Find(&types.UnreadFilter{UserID: currentUserID, ChannelID: channelID}) + if u := uu.FindByChannelId(channelID); u != nil { + tcount = u.InThreadCount + } } if lastReadMessageID > 0 { @@ -491,7 +500,7 @@ func (svc message) MarkAsRead(channelID, threadID, lastReadMessageID uint64) (ui return errors.Wrap(err, "unable to record unread messages") }) - return lastReadMessageID, count, errors.Wrap(err, "unable to mark as read") + return lastReadMessageID, count, tcount, errors.Wrap(err, "unable to mark as read") } // React on a message with an emoji diff --git a/messaging/rest/message.go b/messaging/rest/message.go index 38749d143..6436dbf64 100644 --- a/messaging/rest/message.go +++ b/messaging/rest/message.go @@ -63,11 +63,12 @@ func (ctrl *Message) Delete(ctx context.Context, r *request.MessageDelete) (inte } func (ctrl *Message) MarkAsRead(ctx context.Context, r *request.MessageMarkAsRead) (interface{}, error) { - var messageID, count, err = ctrl.svc.msg.With(ctx).MarkAsRead(r.ChannelID, r.ThreadID, r.LastReadMessageID) + var messageID, count, tcount, err = ctrl.svc.msg.With(ctx).MarkAsRead(r.ChannelID, r.ThreadID, r.LastReadMessageID) return outgoing.Unread{ LastMessageID: messageID, Count: count, + InThreadCount: tcount, }, err } diff --git a/messaging/types/unread.go b/messaging/types/unread.go index 5f612423a..162bbc167 100644 --- a/messaging/types/unread.go +++ b/messaging/types/unread.go @@ -13,6 +13,7 @@ type ( UnreadFilter struct { UserID uint64 + ChannelID uint64 ThreadIDs []uint64 } )