Badge count fix for push notifications when CRT is enabled (#20898)
* Fix * Fixed other cases and tests Co-authored-by: Mattermod <mattermod@users.noreply.github.com>
Этот коммит содержится в:
коммит произвёл
GitHub
родитель
8edd351f27
Коммит
203c6a5013
@@ -219,30 +219,42 @@ func (a *App) getPushNotificationMessage(contentsConfig, postMessage string, exp
|
||||
return senderName + userLocale("api.post.send_notifications_and_forget.push_general_message")
|
||||
}
|
||||
|
||||
func (a *App) getUserBadgeCount(userID string, isCRTEnabled bool) (int, *model.AppError) {
|
||||
unreadCount, err := a.Srv().Store.User().GetUnreadCount(userID, isCRTEnabled)
|
||||
if err != nil {
|
||||
return 0, model.NewAppError("getUserBadgeCount", "app.user.get_unread_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
badgeCount := int(unreadCount)
|
||||
|
||||
if isCRTEnabled {
|
||||
threadUnreadMentions, err := a.Srv().Store.Thread().GetTotalUnreadMentions(userID, "", model.GetUserThreadsOpts{})
|
||||
if err != nil {
|
||||
return 0, model.NewAppError("getUserBadgeCount", "app.user.get_thread_count_for_user.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
badgeCount += int(threadUnreadMentions)
|
||||
}
|
||||
|
||||
return badgeCount, nil
|
||||
}
|
||||
|
||||
func (a *App) clearPushNotificationSync(c request.CTX, currentSessionId, userID, channelID, rootID string) *model.AppError {
|
||||
isCRTEnabled := a.IsCRTEnabledForUser(c, userID)
|
||||
|
||||
badgeCount, err := a.getUserBadgeCount(userID, isCRTEnabled)
|
||||
if err != nil {
|
||||
return model.NewAppError("clearPushNotificationSync", "app.user.get_badge_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
|
||||
msg := &model.PushNotification{
|
||||
Type: model.PushTypeClear,
|
||||
Version: model.PushMessageV2,
|
||||
ChannelId: channelID,
|
||||
RootId: rootID,
|
||||
ContentAvailable: 1,
|
||||
Badge: 0,
|
||||
IsCRTEnabled: a.IsCRTEnabledForUser(c, userID),
|
||||
Badge: badgeCount,
|
||||
IsCRTEnabled: isCRTEnabled,
|
||||
}
|
||||
|
||||
unreadCount, err := a.Srv().Store.User().GetUnreadCount(userID)
|
||||
if err != nil {
|
||||
return model.NewAppError("clearPushNotificationSync", "app.user.get_unread_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
msg.Badge = int(unreadCount)
|
||||
|
||||
if msg.IsCRTEnabled {
|
||||
totalUnreadMentions, err := a.Srv().Store.Thread().GetTotalUnreadMentions(userID, "", model.GetUserThreadsOpts{})
|
||||
if err != nil {
|
||||
return model.NewAppError("clearPushNotificationSync", "app.user.get_thread_count_for_user.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
msg.Badge += int(totalUnreadMentions)
|
||||
}
|
||||
return a.sendPushNotificationToAllSessions(msg, userID, currentSessionId)
|
||||
}
|
||||
|
||||
@@ -260,21 +272,19 @@ func (a *App) clearPushNotification(currentSessionId, userID, channelID, rootID
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) updateMobileAppBadgeSync(userID string) *model.AppError {
|
||||
func (a *App) updateMobileAppBadgeSync(c request.CTX, userID string) *model.AppError {
|
||||
badgeCount, err := a.getUserBadgeCount(userID, a.IsCRTEnabledForUser(c, userID))
|
||||
if err != nil {
|
||||
return model.NewAppError("updateMobileAppBadgeSync", "app.user.get_badge_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
|
||||
msg := &model.PushNotification{
|
||||
Type: model.PushTypeUpdateBadge,
|
||||
Version: model.PushMessageV2,
|
||||
Sound: "none",
|
||||
ContentAvailable: 1,
|
||||
Badge: badgeCount,
|
||||
}
|
||||
|
||||
unreadCount, err := a.Srv().Store.User().GetUnreadCount(userID)
|
||||
if err != nil {
|
||||
return model.NewAppError("updateMobileAppBadgeSync", "app.user.get_unread_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
|
||||
msg.Badge = int(unreadCount)
|
||||
|
||||
return a.sendPushNotificationToAllSessions(msg, userID, "")
|
||||
}
|
||||
|
||||
@@ -345,7 +355,7 @@ func (hub *PushNotificationsHub) start(c request.CTX) {
|
||||
notification.replyToThreadType,
|
||||
)
|
||||
case notificationTypeUpdateBadge:
|
||||
err = hub.app.updateMobileAppBadgeSync(notification.userID)
|
||||
err = hub.app.updateMobileAppBadgeSync(c, notification.userID)
|
||||
default:
|
||||
mlog.Debug("Invalid notification type", mlog.String("notification_type", string(notification.notificationType)))
|
||||
}
|
||||
@@ -566,11 +576,12 @@ func (a *App) BuildPushNotificationMessage(c request.CTX, contentsConfig string,
|
||||
msg = a.buildFullPushNotificationMessage(c, contentsConfig, post, user, channel, channelName, senderName, explicitMention, channelWideMention, replyToThreadType)
|
||||
}
|
||||
|
||||
unreadCount, err := a.Srv().Store.User().GetUnreadCount(user.Id)
|
||||
badgeCount, err := a.getUserBadgeCount(user.Id, a.IsCRTEnabledForUser(c, user.Id))
|
||||
if err != nil {
|
||||
return nil, model.NewAppError("BuildPushNotificationMessage", "app.user.get_unread_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
return nil, model.NewAppError("BuildPushNotificationMessage", "app.user.get_badge_count.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
|
||||
}
|
||||
msg.Badge = int(unreadCount)
|
||||
|
||||
msg.Badge = badgeCount
|
||||
|
||||
return msg, nil
|
||||
}
|
||||
|
||||
@@ -1136,7 +1136,7 @@ func TestClearPushNotificationSync(t *testing.T) {
|
||||
mockStore := th.App.Srv().Store.(*mocks.Store)
|
||||
mockUserStore := mocks.UserStore{}
|
||||
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string")).Return(int64(1), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string"), mock.AnythingOfType("bool")).Return(int64(1), nil)
|
||||
mockPostStore := mocks.PostStore{}
|
||||
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
||||
mockSystemStore := mocks.SystemStore{}
|
||||
@@ -1212,7 +1212,7 @@ func TestUpdateMobileAppBadgeSync(t *testing.T) {
|
||||
mockStore := th.App.Srv().Store.(*mocks.Store)
|
||||
mockUserStore := mocks.UserStore{}
|
||||
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string")).Return(int64(1), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string"), mock.AnythingOfType("bool")).Return(int64(1), nil)
|
||||
mockPostStore := mocks.PostStore{}
|
||||
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
||||
mockSystemStore := mocks.SystemStore{}
|
||||
@@ -1231,9 +1231,10 @@ func TestUpdateMobileAppBadgeSync(t *testing.T) {
|
||||
|
||||
th.App.UpdateConfig(func(cfg *model.Config) {
|
||||
*cfg.EmailSettings.PushNotificationServer = pushServer.URL
|
||||
*cfg.ServiceSettings.CollapsedThreads = model.CollapsedThreadsDisabled
|
||||
})
|
||||
|
||||
err := th.App.updateMobileAppBadgeSync("user1")
|
||||
err := th.App.updateMobileAppBadgeSync(th.Context, "user1")
|
||||
require.Nil(t, err)
|
||||
// Server side verification.
|
||||
// We verify that 2 requests have been sent, and also check the message contents.
|
||||
@@ -1529,7 +1530,7 @@ func BenchmarkPushNotificationThroughput(b *testing.B) {
|
||||
mockStore := th.App.Srv().Store.(*mocks.Store)
|
||||
mockUserStore := mocks.UserStore{}
|
||||
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string")).Return(int64(1), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string"), mock.AnythingOfType("bool")).Return(int64(1), nil)
|
||||
mockPostStore := mocks.PostStore{}
|
||||
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
||||
mockSystemStore := mocks.SystemStore{}
|
||||
|
||||
@@ -135,7 +135,7 @@ func TestHubSessionRevokeRace(t *testing.T) {
|
||||
|
||||
mockUserStore := mocks.UserStore{}
|
||||
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string")).Return(int64(1), nil)
|
||||
mockUserStore.On("GetUnreadCount", mock.AnythingOfType("string"), mock.AnythingOfType("bool")).Return(int64(1), nil)
|
||||
mockPostStore := mocks.PostStore{}
|
||||
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
||||
mockSystemStore := mocks.SystemStore{}
|
||||
|
||||
@@ -6591,6 +6591,10 @@
|
||||
"id": "app.user.get.app_error",
|
||||
"translation": "We encountered an error finding the account."
|
||||
},
|
||||
{
|
||||
"id": "app.user.get_badge_count.app_error",
|
||||
"translation": "We could not get the badge count for the user."
|
||||
},
|
||||
{
|
||||
"id": "app.user.get_by_auth.missing_account.app_error",
|
||||
"translation": "Unable to find an existing account matching your authentication type for this team. This team may require an invite from the team owner to join."
|
||||
|
||||
@@ -11089,7 +11089,7 @@ func (s *OpenTracingLayerUserStore) GetTeamGroupUsers(teamID string) ([]*model.U
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *OpenTracingLayerUserStore) GetUnreadCount(userID string) (int64, error) {
|
||||
func (s *OpenTracingLayerUserStore) GetUnreadCount(userID string, isCRTEnabled bool) (int64, error) {
|
||||
origCtx := s.Root.Store.Context()
|
||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetUnreadCount")
|
||||
s.Root.Store.SetContext(newCtx)
|
||||
@@ -11098,7 +11098,7 @@ func (s *OpenTracingLayerUserStore) GetUnreadCount(userID string) (int64, error)
|
||||
}()
|
||||
|
||||
defer span.Finish()
|
||||
result, err := s.UserStore.GetUnreadCount(userID)
|
||||
result, err := s.UserStore.GetUnreadCount(userID, isCRTEnabled)
|
||||
if err != nil {
|
||||
span.LogFields(spanlog.Error(err))
|
||||
ext.Error.Set(span, true)
|
||||
|
||||
@@ -12660,11 +12660,11 @@ func (s *RetryLayerUserStore) GetTeamGroupUsers(teamID string) ([]*model.User, e
|
||||
|
||||
}
|
||||
|
||||
func (s *RetryLayerUserStore) GetUnreadCount(userID string) (int64, error) {
|
||||
func (s *RetryLayerUserStore) GetUnreadCount(userID string, isCRTEnabled bool) (int64, error) {
|
||||
|
||||
tries := 0
|
||||
for {
|
||||
result, err := s.UserStore.GetUnreadCount(userID)
|
||||
result, err := s.UserStore.GetUnreadCount(userID, isCRTEnabled)
|
||||
if err == nil {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
@@ -1371,9 +1371,18 @@ func (us SqlUserStore) AnalyticsActiveCountForPeriod(startTime int64, endTime in
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func (us SqlUserStore) GetUnreadCount(userId string) (int64, error) {
|
||||
func (us SqlUserStore) GetUnreadCount(userId string, isCRTEnabled bool) (int64, error) {
|
||||
var totalMsgCountColumn = "c.TotalMsgCount"
|
||||
var msgCountColumn = "cm.MsgCount"
|
||||
var mentionCountColumn = "cm.MentionCount"
|
||||
if isCRTEnabled {
|
||||
totalMsgCountColumn = "c.TotalMsgCountRoot"
|
||||
msgCountColumn = "cm.MsgCountRoot"
|
||||
mentionCountColumn = "cm.MentionCountRoot"
|
||||
}
|
||||
|
||||
query := `
|
||||
SELECT SUM(CASE WHEN c.Type = ? THEN (c.TotalMsgCount - cm.MsgCount) ELSE cm.MentionCount END)
|
||||
SELECT SUM(CASE WHEN c.Type = ? THEN (` + totalMsgCountColumn + ` - ` + msgCountColumn + `) ELSE ` + mentionCountColumn + ` END)
|
||||
FROM Channels c
|
||||
INNER JOIN ChannelMembers cm
|
||||
ON cm.ChannelId = c.Id
|
||||
|
||||
@@ -447,7 +447,7 @@ type UserStore interface {
|
||||
PermanentDelete(userID string) error
|
||||
AnalyticsActiveCount(timestamp int64, options model.UserCountOptions) (int64, error)
|
||||
AnalyticsActiveCountForPeriod(startTime int64, endTime int64, options model.UserCountOptions) (int64, error)
|
||||
GetUnreadCount(userID string) (int64, error)
|
||||
GetUnreadCount(userID string, isCRTEnabled bool) (int64, error)
|
||||
GetUnreadCountForChannel(userID string, channelID string) (int64, error)
|
||||
GetAnyUnreadPostCountForChannel(userID string, channelID string) (int64, error)
|
||||
GetRecentlyActiveUsersForTeam(teamID string, offset, limit int, viewRestrictions *model.ViewUsersRestrictions) ([]*model.User, error)
|
||||
|
||||
@@ -960,20 +960,20 @@ func (_m *UserStore) GetTeamGroupUsers(teamID string) ([]*model.User, error) {
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
// GetUnreadCount provides a mock function with given fields: userID
|
||||
func (_m *UserStore) GetUnreadCount(userID string) (int64, error) {
|
||||
ret := _m.Called(userID)
|
||||
// GetUnreadCount provides a mock function with given fields: userID, isCRTEnabled
|
||||
func (_m *UserStore) GetUnreadCount(userID string, isCRTEnabled bool) (int64, error) {
|
||||
ret := _m.Called(userID, isCRTEnabled)
|
||||
|
||||
var r0 int64
|
||||
if rf, ok := ret.Get(0).(func(string) int64); ok {
|
||||
r0 = rf(userID)
|
||||
if rf, ok := ret.Get(0).(func(string, bool) int64); ok {
|
||||
r0 = rf(userID, isCRTEnabled)
|
||||
} else {
|
||||
r0 = ret.Get(0).(int64)
|
||||
}
|
||||
|
||||
var r1 error
|
||||
if rf, ok := ret.Get(1).(func(string) error); ok {
|
||||
r1 = rf(userID)
|
||||
if rf, ok := ret.Get(1).(func(string, bool) error); ok {
|
||||
r1 = rf(userID, isCRTEnabled)
|
||||
} else {
|
||||
r1 = ret.Error(1)
|
||||
}
|
||||
|
||||
@@ -2440,14 +2440,23 @@ func testUserUnreadCount(t *testing.T, ss store.Store) {
|
||||
nErr = ss.Channel().IncrementMentionCount(c2.Id, []string{u2.Id}, false)
|
||||
require.NoError(t, nErr)
|
||||
|
||||
badge, unreadCountErr := ss.User().GetUnreadCount(u2.Id)
|
||||
badge, unreadCountErr := ss.User().GetUnreadCount(u2.Id, false)
|
||||
require.NoError(t, unreadCountErr)
|
||||
require.Equal(t, int64(3), badge, "should have 3 unread messages")
|
||||
|
||||
badge, unreadCountErr = ss.User().GetUnreadCount(u3.Id)
|
||||
badge, unreadCountErr = ss.User().GetUnreadCount(u3.Id, false)
|
||||
require.NoError(t, unreadCountErr)
|
||||
require.Equal(t, int64(1), badge, "should have 1 unread message")
|
||||
|
||||
// Increment root mentions by 1
|
||||
nErr = ss.Channel().IncrementMentionCount(c1.Id, []string{u3.Id}, true)
|
||||
require.NoError(t, nErr)
|
||||
|
||||
// CRT is enabled, only root mentions are counted
|
||||
badge, unreadCountErr = ss.User().GetUnreadCount(u3.Id, true)
|
||||
require.NoError(t, unreadCountErr)
|
||||
require.Equal(t, int64(1), badge, "should have 1 unread message with CRT")
|
||||
|
||||
badge, unreadCountErr = ss.User().GetUnreadCountForChannel(u2.Id, c1.Id)
|
||||
require.NoError(t, unreadCountErr)
|
||||
require.Equal(t, int64(1), badge, "should have 1 unread messages for that channel")
|
||||
|
||||
@@ -9983,10 +9983,10 @@ func (s *TimerLayerUserStore) GetTeamGroupUsers(teamID string) ([]*model.User, e
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *TimerLayerUserStore) GetUnreadCount(userID string) (int64, error) {
|
||||
func (s *TimerLayerUserStore) GetUnreadCount(userID string, isCRTEnabled bool) (int64, error) {
|
||||
start := time.Now()
|
||||
|
||||
result, err := s.UserStore.GetUnreadCount(userID)
|
||||
result, err := s.UserStore.GetUnreadCount(userID, isCRTEnabled)
|
||||
|
||||
elapsed := float64(time.Since(start)) / float64(time.Second)
|
||||
if s.Root.Metrics != nil {
|
||||
|
||||
Ссылка в новой задаче
Block a user