From 8844141df39da9e254cd6464231bf2b81ed6e1a6 Mon Sep 17 00:00:00 2001 From: Eli Yukelzon Date: Thu, 15 Oct 2020 18:01:16 +0300 Subject: [PATCH] MM-28249 Auto follow threads (#15878) Co-authored-by: Mattermod --- app/app_iface.go | 1 + app/notification.go | 36 ++++++- app/opentracing/opentracing_layer.go | 22 ++++ app/post.go | 10 ++ app/post_test.go | 57 ++++++++++ i18n/en.json | 4 + model/config.go | 6 ++ model/thread.go | 13 +++ store/opentracinglayer/opentracinglayer.go | 108 +++++++++++++++++++ store/retrylayer/retrylayer.go | 120 +++++++++++++++++++++ store/sqlstore/post_store.go | 8 +- store/sqlstore/thread_store.go | 88 ++++++++++++++- store/store.go | 7 ++ store/storetest/mocks/ThreadStore.go | 120 +++++++++++++++++++++ store/timerlayer/timerlayer.go | 96 +++++++++++++++++ 15 files changed, 686 insertions(+), 10 deletions(-) diff --git a/app/app_iface.go b/app/app_iface.go index 12dd59bfa5..7302cc69a8 100644 --- a/app/app_iface.go +++ b/app/app_iface.go @@ -685,6 +685,7 @@ type AppIface interface { GetTeamsForUser(userId string) ([]*model.Team, *model.AppError) GetTeamsUnreadForUser(excludeTeamId string, userId string) ([]*model.TeamUnread, *model.AppError) GetTermsOfService(id string) (*model.TermsOfService, *model.AppError) + GetThreadMembershipsForUser(userId string) ([]*model.ThreadMembership, error) GetUploadSession(uploadId string) (*model.UploadSession, *model.AppError) GetUploadSessionsForUser(userId string) ([]*model.UploadSession, *model.AppError) GetUser(userId string) (*model.User, *model.AppError) diff --git a/app/notification.go b/app/notification.go index e45add5a5b..eca1ae2fa2 100644 --- a/app/notification.go +++ b/app/notification.go @@ -159,20 +159,37 @@ func (a *App) SendNotifications(post *model.Post, team *model.Team, channel *mod mentionedUsersList := make([]string, 0, len(mentions.Mentions)) updateMentionChans := []chan *model.AppError{} + mentionAutofollowChans := []chan *model.AppError{} + + // for each mention, make sure to update thread autofollow + for id := range mentions.Mentions { + mac := make(chan *model.AppError, 1) + go func(userId string) { + defer close(mac) + if *a.Config().ServiceSettings.ThreadAutoFollow && post.RootId != "" { + nErr := a.Srv().Store.Thread().CreateMembershipIfNeeded(userId, post.RootId) + if nErr != nil { + mac <- model.NewAppError("SendNotifications", "app.channel.autofollow.app_error", nil, nErr.Error(), http.StatusInternalServerError) + return + } + } + mac <- nil + }(id) + mentionAutofollowChans = append(mentionAutofollowChans, mac) + } for id := range mentions.Mentions { mentionedUsersList = append(mentionedUsersList, id) umc := make(chan *model.AppError, 1) go func(userId string) { + defer close(umc) nErr := a.Srv().Store.Channel().IncrementMentionCount(post.ChannelId, userId) if nErr != nil { umc <- model.NewAppError("SendNotifications", "app.channel.increment_mention_count.app_error", nil, nErr.Error(), http.StatusInternalServerError) - } else { - umc <- nil + return } - - close(umc) + umc <- nil }(id) updateMentionChans = append(updateMentionChans, umc) } @@ -254,6 +271,17 @@ func (a *App) SendNotifications(post *model.Post, team *model.Team, channel *mod } } + // Log the problems that might have occurred while auto following the thread + for _, mac := range mentionAutofollowChans { + if err := <-mac; err != nil { + mlog.Warn( + "Failed to update thread autofollow from mention", + mlog.String("post_id", post.Id), + mlog.String("channel_id", post.ChannelId), + mlog.Err(err), + ) + } + } sendPushNotifications := false if *a.Config().EmailSettings.SendPushNotifications { pushServer := *a.Config().EmailSettings.PushNotificationServer diff --git a/app/opentracing/opentracing_layer.go b/app/opentracing/opentracing_layer.go index 64a28c3d58..23b18d10b2 100644 --- a/app/opentracing/opentracing_layer.go +++ b/app/opentracing/opentracing_layer.go @@ -8397,6 +8397,28 @@ func (a *OpenTracingAppLayer) GetTermsOfService(id string) (*model.TermsOfServic return resultVar0, resultVar1 } +func (a *OpenTracingAppLayer) GetThreadMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + origCtx := a.ctx + span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetThreadMembershipsForUser") + + a.ctx = newCtx + a.app.Srv().Store.SetContext(newCtx) + defer func() { + a.app.Srv().Store.SetContext(origCtx) + a.ctx = origCtx + }() + + defer span.Finish() + resultVar0, resultVar1 := a.app.GetThreadMembershipsForUser(userId) + + if resultVar1 != nil { + span.LogFields(spanlog.Error(resultVar1)) + ext.Error.Set(span, true) + } + + return resultVar0, resultVar1 +} + func (a *OpenTracingAppLayer) GetTotalUsersStats(viewRestrictions *model.ViewUsersRestrictions) (*model.UsersStats, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetTotalUsersStats") diff --git a/app/post.go b/app/post.go index a15b215264..da731f8f2a 100644 --- a/app/post.go +++ b/app/post.go @@ -443,6 +443,12 @@ func (a *App) handlePostEvents(post *model.Post, user *model.User, channel *mode return err } + if *a.Config().ServiceSettings.ThreadAutoFollow && post.RootId != "" { + if err := a.Srv().Store.Thread().CreateMembershipIfNeeded(post.UserId, post.RootId); err != nil { + return err + } + } + if post.Type != model.POST_AUTO_RESPONDER { // don't respond to an auto-responder a.Srv().Go(func() { _, err := a.SendAutoResponseIfNecessary(channel, user) @@ -1450,3 +1456,7 @@ func isPostMention(user *model.User, post *model.Post, keywords map[string][]str return false } + +func (a *App) GetThreadMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + return a.Srv().Store.Thread().GetMembershipsForUser(userId) +} diff --git a/app/post_test.go b/app/post_test.go index b344e574cc..c3abf9bd7e 100644 --- a/app/post_test.go +++ b/app/post_test.go @@ -1846,3 +1846,60 @@ func TestFillInPostProps(t *testing.T) { assert.Equal(t, post1.Props, model.StringInterface{"disable_group_highlight": true}) }) } + +func TestThreadMembership(t *testing.T) { + t.Run("should update memberships for conversation participants", func(t *testing.T) { + th := Setup(t).InitBasic() + defer th.TearDown() + + user1 := th.BasicUser + user2 := th.BasicUser2 + + channel := th.CreateChannel(th.BasicTeam) + th.AddUserToChannel(user2, channel) + + postRoot, err := th.App.CreatePost(&model.Post{ + UserId: user1.Id, + ChannelId: channel.Id, + Message: "root post", + }, channel, false, true) + require.Nil(t, err) + + _, err = th.App.CreatePost(&model.Post{ + UserId: user1.Id, + ChannelId: channel.Id, + RootId: postRoot.Id, + Message: fmt.Sprintf("@%s", user2.Username), + }, channel, false, true) + require.Nil(t, err) + + // first user should now be part of the thread since they replied to a post + memberships, err2 := th.App.GetThreadMembershipsForUser(user1.Id) + require.Nil(t, err2) + require.Len(t, memberships, 1) + // second user should also be part of a thread since they were mentioned + memberships, err2 = th.App.GetThreadMembershipsForUser(user2.Id) + require.Nil(t, err2) + require.Len(t, memberships, 1) + + post2, err := th.App.CreatePost(&model.Post{ + UserId: user2.Id, + ChannelId: channel.Id, + Message: "second post", + }, channel, false, true) + require.Nil(t, err) + + _, err = th.App.CreatePost(&model.Post{ + UserId: user2.Id, + ChannelId: channel.Id, + RootId: post2.Id, + Message: fmt.Sprintf("@%s", user1.Username), + }, channel, false, true) + require.Nil(t, err) + + // first user should now be part of two threads + memberships, err2 = th.App.GetThreadMembershipsForUser(user1.Id) + require.Nil(t, err2) + require.Len(t, memberships, 2) + }) +} diff --git a/i18n/en.json b/i18n/en.json index 187933ef2c..dd2e576174 100644 --- a/i18n/en.json +++ b/i18n/en.json @@ -3502,6 +3502,10 @@ "id": "app.channel.analytics_type_count.app_error", "translation": "Unable to get channel type counts." }, + { + "id": "app.channel.autofollow.app_error", + "translation": "Failed to update thread membership for mentioned user" + }, { "id": "app.channel.clear_all_custom_role_assignments.select.app_error", "translation": "Failed to retrieve the channel members." diff --git a/model/config.go b/model/config.go index 7f6b2b6e35..bae125e618 100644 --- a/model/config.go +++ b/model/config.go @@ -347,6 +347,7 @@ type ServiceSettings struct { EnableLocalMode *bool LocalModeSocketLocation *string EnableAWSMetering *bool + ThreadAutoFollow *bool `access:"experimental"` } func (s *ServiceSettings) SetDefaults(isUpdate bool) { @@ -764,6 +765,11 @@ func (s *ServiceSettings) SetDefaults(isUpdate bool) { if s.EnableAWSMetering == nil { s.EnableAWSMetering = NewBool(false) } + + if s.ThreadAutoFollow == nil { + s.ThreadAutoFollow = NewBool(true) + } + } type ClusterSettings struct { diff --git a/model/thread.go b/model/thread.go index 76707c5f19..8485032b2c 100644 --- a/model/thread.go +++ b/model/thread.go @@ -22,3 +22,16 @@ func (o *Thread) ToJson() string { func (o *Thread) Etag() string { return Etag(o.PostId, o.LastReplyAt) } + +type ThreadMembership struct { + PostId string `json:"post_id"` + UserId string `json:"user_id"` + Following bool `json:"following"` + LastViewed int64 `json:"last_view_at"` + LastUpdated int64 `json:"last_update_at"` +} + +func (o *ThreadMembership) ToJson() string { + b, _ := json.Marshal(o) + return string(b) +} diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index 059628e9ef..54d8859066 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -7612,6 +7612,24 @@ func (s *OpenTracingLayerTermsOfServiceStore) Save(termsOfService *model.TermsOf return result, err } +func (s *OpenTracingLayerThreadStore) CreateMembershipIfNeeded(userId string, postId string) error { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.CreateMembershipIfNeeded") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + err := s.ThreadStore.CreateMembershipIfNeeded(userId, postId) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return err +} + func (s *OpenTracingLayerThreadStore) Delete(postId string) error { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.Delete") @@ -7630,6 +7648,24 @@ func (s *OpenTracingLayerThreadStore) Delete(postId string) error { return err } +func (s *OpenTracingLayerThreadStore) DeleteMembershipForUser(userId string, postId string) error { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.DeleteMembershipForUser") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + err := s.ThreadStore.DeleteMembershipForUser(userId, postId) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return err +} + func (s *OpenTracingLayerThreadStore) Get(id string) (*model.Thread, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.Get") @@ -7648,6 +7684,42 @@ func (s *OpenTracingLayerThreadStore) Get(id string) (*model.Thread, error) { return result, err } +func (s *OpenTracingLayerThreadStore) GetMembershipForUser(userId string, postId string) (*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetMembershipForUser") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.GetMembershipForUser(userId, postId) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + +func (s *OpenTracingLayerThreadStore) GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetMembershipsForUser") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.GetMembershipsForUser(userId) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerThreadStore) Save(thread *model.Thread) (*model.Thread, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.Save") @@ -7666,6 +7738,24 @@ func (s *OpenTracingLayerThreadStore) Save(thread *model.Thread) (*model.Thread, return result, err } +func (s *OpenTracingLayerThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.SaveMembership") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.SaveMembership(membership) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerThreadStore) SaveMultiple(thread []*model.Thread) ([]*model.Thread, int, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.SaveMultiple") @@ -7702,6 +7792,24 @@ func (s *OpenTracingLayerThreadStore) Update(thread *model.Thread) (*model.Threa return result, err } +func (s *OpenTracingLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.UpdateMembership") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.UpdateMembership(membership) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerTokenStore) Cleanup() { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TokenStore.Cleanup") diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index da1668f69e..e4fa2d48e9 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -7642,6 +7642,26 @@ func (s *RetryLayerTermsOfServiceStore) Save(termsOfService *model.TermsOfServic } +func (s *RetryLayerThreadStore) CreateMembershipIfNeeded(userId string, postId string) error { + + tries := 0 + for { + err := s.ThreadStore.CreateMembershipIfNeeded(userId, postId) + if err == nil { + return nil + } + if !isRepeatableError(err) { + return err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return err + } + } + +} + func (s *RetryLayerThreadStore) Delete(postId string) error { tries := 0 @@ -7662,6 +7682,26 @@ func (s *RetryLayerThreadStore) Delete(postId string) error { } +func (s *RetryLayerThreadStore) DeleteMembershipForUser(userId string, postId string) error { + + tries := 0 + for { + err := s.ThreadStore.DeleteMembershipForUser(userId, postId) + if err == nil { + return nil + } + if !isRepeatableError(err) { + return err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return err + } + } + +} + func (s *RetryLayerThreadStore) Get(id string) (*model.Thread, error) { tries := 0 @@ -7682,6 +7722,46 @@ func (s *RetryLayerThreadStore) Get(id string) (*model.Thread, error) { } +func (s *RetryLayerThreadStore) GetMembershipForUser(userId string, postId string) (*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.GetMembershipForUser(userId, postId) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + } + +} + +func (s *RetryLayerThreadStore) GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.GetMembershipsForUser(userId) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + } + +} + func (s *RetryLayerThreadStore) Save(thread *model.Thread) (*model.Thread, error) { tries := 0 @@ -7702,6 +7782,26 @@ func (s *RetryLayerThreadStore) Save(thread *model.Thread) (*model.Thread, error } +func (s *RetryLayerThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.SaveMembership(membership) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + } + +} + func (s *RetryLayerThreadStore) SaveMultiple(thread []*model.Thread) ([]*model.Thread, int, error) { tries := 0 @@ -7742,6 +7842,26 @@ func (s *RetryLayerThreadStore) Update(thread *model.Thread) (*model.Thread, err } +func (s *RetryLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.UpdateMembership(membership) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + } + +} + func (s *RetryLayerTokenStore) Cleanup() { s.TokenStore.Cleanup() diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index c4a6a8ce5b..865734bb32 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -1937,9 +1937,11 @@ func (s *SqlPostStore) cleanupThreads(postId, rootId, userId string) error { } } } - _, err := s.GetMaster().Exec("DELETE FROM Threads WHERE PostId = :Id", map[string]interface{}{"Id": postId}) - if err != nil { - return errors.Wrap(err, "failed to update Threads") + if _, err := s.GetMaster().Exec("DELETE FROM Threads WHERE PostId = :Id", map[string]interface{}{"Id": postId}); err != nil { + return errors.Wrap(err, "failed to delete Threads") + } + if _, err := s.GetMaster().Exec("DELETE FROM ThreadMemberships WHERE PostId = :Id", map[string]interface{}{"Id": postId}); err != nil { + return errors.Wrap(err, "failed to delete ThreadMemberships") } return nil } diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index 3254076a7a..92ffaff7ef 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -7,7 +7,9 @@ import ( "database/sql" "github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/store" + "github.com/mattermost/mattermost-server/v5/utils" "github.com/pkg/errors" + "time" sq "github.com/Masterminds/squirrel" ) @@ -25,9 +27,12 @@ func newSqlThreadStore(sqlStore SqlStore) store.ThreadStore { } for _, db := range sqlStore.GetAllConns() { - table := db.AddTableWithName(model.Thread{}, "Threads").SetKeys(false, "PostId") - table.ColMap("PostId").SetMaxSize(26) - table.ColMap("Participants").SetMaxSize(0) + tableThreads := db.AddTableWithName(model.Thread{}, "Threads").SetKeys(false, "PostId") + tableThreads.ColMap("PostId").SetMaxSize(26) + tableThreads.ColMap("Participants").SetMaxSize(0) + tableThreadMemberships := db.AddTableWithName(model.ThreadMembership{}, "ThreadMemberships").SetKeys(false, "PostId", "UserId") + tableThreadMemberships.ColMap("PostId").SetMaxSize(26) + tableThreadMemberships.ColMap("UserId").SetMaxSize(26) } return s @@ -49,6 +54,11 @@ func threadToSlice(thread *model.Thread) []interface{} { func (s *SqlThreadStore) createIndexesIfNotExists() { s.CreateIndexIfNotExists("idx_threads_last_reply_at", "Threads", "LastReplyAt") s.CreateIndexIfNotExists("idx_threads_post_id", "Threads", "PostId") + + s.CreateIndexIfNotExists("idx_thread_memberships_last_update_at", "ThreadMemberships", "LastUpdated") + s.CreateIndexIfNotExists("idx_thread_memberships_last_view_at", "ThreadMemberships", "LastViewed") + s.CreateIndexIfNotExists("idx_thread_memberships_post_id", "ThreadMemberships", "PostId") + s.CreateIndexIfNotExists("idx_thread_memberships_user_id", "ThreadMemberships", "UserId") } func (s *SqlThreadStore) SaveMultiple(threads []*model.Thread) ([]*model.Thread, int, error) { @@ -106,3 +116,75 @@ func (s *SqlThreadStore) Delete(threadId string) error { return nil } + +func (s *SqlThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + if err := s.GetMaster().Insert(membership); err != nil { + return nil, errors.Wrapf(err, "failed to save thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) + } + + return membership, nil +} + +func (s *SqlThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + if _, err := s.GetMaster().Update(membership); err != nil { + return nil, errors.Wrapf(err, "failed to update thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) + } + + return membership, nil +} + +func (s *SqlThreadStore) GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + var memberships []*model.ThreadMembership + _, err := s.GetReplica().Select(&memberships, "SELECT * from ThreadMemberships WHERE UserId = :UserId", map[string]interface{}{"UserId": userId}) + if err != nil { + return nil, errors.Wrapf(err, "failed to get thread membership with userid=%s", userId) + } + return memberships, nil +} + +func (s *SqlThreadStore) GetMembershipForUser(userId, postId string) (*model.ThreadMembership, error) { + var membership model.ThreadMembership + err := s.GetReplica().SelectOne(&membership, "SELECT * from ThreadMemberships WHERE UserId = :UserId AND PostId = :PostId", map[string]interface{}{"UserId": userId, "PostId": postId}) + if err != nil { + if err == sql.ErrNoRows { + return nil, store.NewErrNotFound("Thread", postId) + } + return nil, errors.Wrapf(err, "failed to get thread membership with userid=%s postid=%s", userId, postId) + } + return &membership, nil +} + +func (s *SqlThreadStore) DeleteMembershipForUser(userId string, postId string) error { + if _, err := s.GetMaster().Exec("DELETE FROM ThreadMemberships Where PostId = :PostId AND UserId = :UserId", map[string]interface{}{"PostId": postId, "UserId": userId}); err != nil { + return errors.Wrap(err, "failed to update thread membership") + } + + return nil +} + +func (s *SqlThreadStore) CreateMembershipIfNeeded(userId, postId string) error { + membership, err := s.GetMembershipForUser(userId, postId) + now := utils.MillisFromTime(time.Now()) + if err == nil { + if !membership.Following { + membership.Following = true + membership.LastUpdated = now + _, err = s.UpdateMembership(membership) + } + return err + } + + var nfErr *store.ErrNotFound + + if !errors.As(err, &nfErr) { + return errors.Wrap(err, "failed to get thread membership") + } + _, err = s.SaveMembership(&model.ThreadMembership{ + PostId: postId, + UserId: userId, + Following: true, + LastViewed: 0, + LastUpdated: now, + }) + return err +} diff --git a/store/store.go b/store/store.go index a73ba48d53..4819cf702a 100644 --- a/store/store.go +++ b/store/store.go @@ -252,6 +252,13 @@ type ThreadStore interface { Update(thread *model.Thread) (*model.Thread, error) Get(id string) (*model.Thread, error) Delete(postId string) error + + SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) + UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) + GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) + GetMembershipForUser(userId, postId string) (*model.ThreadMembership, error) + DeleteMembershipForUser(userId, postId string) error + CreateMembershipIfNeeded(userId, postId string) error } type PostStore interface { diff --git a/store/storetest/mocks/ThreadStore.go b/store/storetest/mocks/ThreadStore.go index 5050c68830..24a367db4f 100644 --- a/store/storetest/mocks/ThreadStore.go +++ b/store/storetest/mocks/ThreadStore.go @@ -14,6 +14,20 @@ type ThreadStore struct { mock.Mock } +// CreateMembershipIfNeeded provides a mock function with given fields: userId, postId +func (_m *ThreadStore) CreateMembershipIfNeeded(userId string, postId string) error { + ret := _m.Called(userId, postId) + + var r0 error + if rf, ok := ret.Get(0).(func(string, string) error); ok { + r0 = rf(userId, postId) + } else { + r0 = ret.Error(0) + } + + return r0 +} + // Delete provides a mock function with given fields: postId func (_m *ThreadStore) Delete(postId string) error { ret := _m.Called(postId) @@ -28,6 +42,20 @@ func (_m *ThreadStore) Delete(postId string) error { return r0 } +// DeleteMembershipForUser provides a mock function with given fields: userId, postId +func (_m *ThreadStore) DeleteMembershipForUser(userId string, postId string) error { + ret := _m.Called(userId, postId) + + var r0 error + if rf, ok := ret.Get(0).(func(string, string) error); ok { + r0 = rf(userId, postId) + } else { + r0 = ret.Error(0) + } + + return r0 +} + // Get provides a mock function with given fields: id func (_m *ThreadStore) Get(id string) (*model.Thread, error) { ret := _m.Called(id) @@ -51,6 +79,52 @@ func (_m *ThreadStore) Get(id string) (*model.Thread, error) { return r0, r1 } +// GetMembershipForUser provides a mock function with given fields: userId, postId +func (_m *ThreadStore) GetMembershipForUser(userId string, postId string) (*model.ThreadMembership, error) { + ret := _m.Called(userId, postId) + + var r0 *model.ThreadMembership + if rf, ok := ret.Get(0).(func(string, string) *model.ThreadMembership); ok { + r0 = rf(userId, postId) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.ThreadMembership) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(string, string) error); ok { + r1 = rf(userId, postId) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// GetMembershipsForUser provides a mock function with given fields: userId +func (_m *ThreadStore) GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + ret := _m.Called(userId) + + var r0 []*model.ThreadMembership + if rf, ok := ret.Get(0).(func(string) []*model.ThreadMembership); ok { + r0 = rf(userId) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.ThreadMembership) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(userId) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Save provides a mock function with given fields: thread func (_m *ThreadStore) Save(thread *model.Thread) (*model.Thread, error) { ret := _m.Called(thread) @@ -74,6 +148,29 @@ func (_m *ThreadStore) Save(thread *model.Thread) (*model.Thread, error) { return r0, r1 } +// SaveMembership provides a mock function with given fields: membership +func (_m *ThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + ret := _m.Called(membership) + + var r0 *model.ThreadMembership + if rf, ok := ret.Get(0).(func(*model.ThreadMembership) *model.ThreadMembership); ok { + r0 = rf(membership) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.ThreadMembership) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(*model.ThreadMembership) error); ok { + r1 = rf(membership) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // SaveMultiple provides a mock function with given fields: thread func (_m *ThreadStore) SaveMultiple(thread []*model.Thread) ([]*model.Thread, int, error) { ret := _m.Called(thread) @@ -126,3 +223,26 @@ func (_m *ThreadStore) Update(thread *model.Thread) (*model.Thread, error) { return r0, r1 } + +// UpdateMembership provides a mock function with given fields: membership +func (_m *ThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + ret := _m.Called(membership) + + var r0 *model.ThreadMembership + if rf, ok := ret.Get(0).(func(*model.ThreadMembership) *model.ThreadMembership); ok { + r0 = rf(membership) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*model.ThreadMembership) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(*model.ThreadMembership) error); ok { + r1 = rf(membership) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index 82d5184581..092a9eae84 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -6872,6 +6872,22 @@ func (s *TimerLayerTermsOfServiceStore) Save(termsOfService *model.TermsOfServic return result, err } +func (s *TimerLayerThreadStore) CreateMembershipIfNeeded(userId string, postId string) error { + start := timemodule.Now() + + err := s.ThreadStore.CreateMembershipIfNeeded(userId, postId) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.CreateMembershipIfNeeded", success, elapsed) + } + return err +} + func (s *TimerLayerThreadStore) Delete(postId string) error { start := timemodule.Now() @@ -6888,6 +6904,22 @@ func (s *TimerLayerThreadStore) Delete(postId string) error { return err } +func (s *TimerLayerThreadStore) DeleteMembershipForUser(userId string, postId string) error { + start := timemodule.Now() + + err := s.ThreadStore.DeleteMembershipForUser(userId, postId) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.DeleteMembershipForUser", success, elapsed) + } + return err +} + func (s *TimerLayerThreadStore) Get(id string) (*model.Thread, error) { start := timemodule.Now() @@ -6904,6 +6936,38 @@ func (s *TimerLayerThreadStore) Get(id string) (*model.Thread, error) { return result, err } +func (s *TimerLayerThreadStore) GetMembershipForUser(userId string, postId string) (*model.ThreadMembership, error) { + start := timemodule.Now() + + result, err := s.ThreadStore.GetMembershipForUser(userId, postId) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.GetMembershipForUser", success, elapsed) + } + return result, err +} + +func (s *TimerLayerThreadStore) GetMembershipsForUser(userId string) ([]*model.ThreadMembership, error) { + start := timemodule.Now() + + result, err := s.ThreadStore.GetMembershipsForUser(userId) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.GetMembershipsForUser", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) Save(thread *model.Thread) (*model.Thread, error) { start := timemodule.Now() @@ -6920,6 +6984,22 @@ func (s *TimerLayerThreadStore) Save(thread *model.Thread) (*model.Thread, error return result, err } +func (s *TimerLayerThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + start := timemodule.Now() + + result, err := s.ThreadStore.SaveMembership(membership) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.SaveMembership", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) SaveMultiple(thread []*model.Thread) ([]*model.Thread, int, error) { start := timemodule.Now() @@ -6952,6 +7032,22 @@ func (s *TimerLayerThreadStore) Update(thread *model.Thread) (*model.Thread, err return result, err } +func (s *TimerLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { + start := timemodule.Now() + + result, err := s.ThreadStore.UpdateMembership(membership) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.UpdateMembership", success, elapsed) + } + return result, err +} + func (s *TimerLayerTokenStore) Cleanup() { start := timemodule.Now()