From 597a2b77cd999f68be49599f0189a0a3e465c111 Mon Sep 17 00:00:00 2001 From: Eli Yukelzon Date: Wed, 5 Feb 2020 13:27:35 +0100 Subject: [PATCH] MM-17468 - Improve thread fetching (#13653) * Revert "Thread fetching revert (#13616)" This reverts commit 8e0fe90897ecc120c3b4da06950d95e18ab07eb4. * renamed query param for clarity Co-authored-by: mattermod --- api4/post.go | 18 +- app/auto_posts.go | 10 +- app/channel_test.go | 2 +- app/command_loadtest.go | 35 ++++ app/file.go | 2 +- app/plugin_api.go | 10 +- app/post.go | 52 +++--- model/post.go | 18 +- store/localcachelayer/main_test.go | 15 +- store/localcachelayer/post_layer.go | 25 +-- store/localcachelayer/post_layer_test.go | 34 ++-- store/sqlstore/post_store.go | 227 ++++++++++++++--------- store/store.go | 10 +- store/storetest/mocks/PostStore.go | 70 +++---- store/storetest/post_store.go | 223 ++++++++++++++++------ store/storetest/reaction_store.go | 76 +++----- store/timer_layer.go | 36 +++- 17 files changed, 541 insertions(+), 322 deletions(-) diff --git a/api4/post.go b/api4/post.go index 6c7bf4d153..d9ef322d52 100644 --- a/api4/post.go +++ b/api4/post.go @@ -148,6 +148,7 @@ func getPostsForChannel(c *Context, w http.ResponseWriter, r *http.Request) { return } } + skipFetchThreads := r.URL.Query().Get("skipFetchThreads") == "true" channelId := c.Params.ChannelId page := c.Params.Page @@ -163,7 +164,7 @@ func getPostsForChannel(c *Context, w http.ResponseWriter, r *http.Request) { etag := "" if since > 0 { - list, err = c.App.GetPostsSince(channelId, since) + list, err = c.App.GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: since, SkipFetchThreads: skipFetchThreads}) } else if len(afterPost) > 0 { etag = c.App.GetPostsEtag(channelId) @@ -171,7 +172,7 @@ func getPostsForChannel(c *Context, w http.ResponseWriter, r *http.Request) { return } - list, err = c.App.GetPostsAfterPost(channelId, afterPost, page, perPage) + list, err = c.App.GetPostsAfterPost(model.GetPostsOptions{ChannelId: channelId, PostId: afterPost, Page: page, PerPage: perPage, SkipFetchThreads: skipFetchThreads}) } else if len(beforePost) > 0 { etag = c.App.GetPostsEtag(channelId) @@ -179,7 +180,7 @@ func getPostsForChannel(c *Context, w http.ResponseWriter, r *http.Request) { return } - list, err = c.App.GetPostsBeforePost(channelId, beforePost, page, perPage) + list, err = c.App.GetPostsBeforePost(model.GetPostsOptions{ChannelId: channelId, PostId: beforePost, Page: page, PerPage: perPage, SkipFetchThreads: skipFetchThreads}) } else { etag = c.App.GetPostsEtag(channelId) @@ -187,7 +188,7 @@ func getPostsForChannel(c *Context, w http.ResponseWriter, r *http.Request) { return } - list, err = c.App.GetPostsPage(channelId, page, perPage) + list, err = c.App.GetPostsPage(model.GetPostsOptions{ChannelId: channelId, Page: page, PerPage: perPage, SkipFetchThreads: skipFetchThreads}) } if err != nil { @@ -223,7 +224,8 @@ func getPostsForChannelAroundLastUnread(c *Context, w http.ResponseWriter, r *ht return } - postList, err := c.App.GetPostsForChannelAroundLastUnread(channelId, userId, c.Params.LimitBefore, c.Params.LimitAfter) + skipFetchThreads := r.URL.Query().Get("skipFetchThreads") == "true" + postList, err := c.App.GetPostsForChannelAroundLastUnread(channelId, userId, c.Params.LimitBefore, c.Params.LimitAfter, skipFetchThreads) if err != nil { c.Err = err return @@ -237,7 +239,7 @@ func getPostsForChannelAroundLastUnread(c *Context, w http.ResponseWriter, r *ht return } - postList, err = c.App.GetPostsPage(channelId, app.PAGE_DEFAULT, c.Params.LimitBefore) + postList, err = c.App.GetPostsPage(model.GetPostsOptions{ChannelId: channelId, Page: app.PAGE_DEFAULT, PerPage: c.Params.LimitBefore, SkipFetchThreads: skipFetchThreads}) if err != nil { c.Err = err return @@ -391,8 +393,8 @@ func getPostThread(c *Context, w http.ResponseWriter, r *http.Request) { if c.Err != nil { return } - - list, err := c.App.GetPostThread(c.Params.PostId) + skipFetchThreads := r.URL.Query().Get("skipFetchThreads") == "true" + list, err := c.App.GetPostThread(c.Params.PostId, skipFetchThreads) if err != nil { c.Err = err return diff --git a/app/auto_posts.go b/app/auto_posts.go index 2cd2c48b55..15a6b76d3b 100644 --- a/app/auto_posts.go +++ b/app/auto_posts.go @@ -66,6 +66,10 @@ func (cfg *AutoPostCreator) UploadTestFile() ([]string, bool) { } func (cfg *AutoPostCreator) CreateRandomPost() (*model.Post, bool) { + return cfg.CreateRandomPostNested("", "") +} + +func (cfg *AutoPostCreator) CreateRandomPostNested(parentId, rootId string) (*model.Post, bool) { var fileIds []string if cfg.HasImage { var err1 bool @@ -84,10 +88,12 @@ func (cfg *AutoPostCreator) CreateRandomPost() (*model.Post, bool) { post := &model.Post{ ChannelId: cfg.channelid, + ParentId: parentId, + RootId: rootId, Message: postText, FileIds: fileIds} - rpost, err2 := cfg.client.CreatePost(post) - if err2 != nil { + rpost, resp := cfg.client.CreatePost(post) + if resp != nil && resp.Error != nil { return nil, false } return rpost, true diff --git a/app/channel_test.go b/app/channel_test.go index 9d39b0a3eb..38ad40cb8e 100644 --- a/app/channel_test.go +++ b/app/channel_test.go @@ -502,7 +502,7 @@ func TestAddChannelMemberNoUserRequestor(t *testing.T) { } assert.Equal(t, groupUserIds, channelMemberHistoryUserIds) - postList, err := th.App.Srv.Store.Post().GetPosts(channel.Id, 0, 1, false) + postList, err := th.App.Srv.Store.Post().GetPosts(model.GetPostsOptions{ChannelId: channel.Id, Page: 0, PerPage: 1}, false) require.Nil(t, err) if assert.Len(t, postList.Order, 1) { diff --git a/app/command_loadtest.go b/app/command_loadtest.go index e0acd8ff9b..d93b5ad7b0 100644 --- a/app/command_loadtest.go +++ b/app/command_loadtest.go @@ -39,6 +39,9 @@ var usage = `Mattermost testing commands to help configure the system Example: /test channels fuzz 5 10 + ThreadedPost - create a large threaded post + /test threaded_post + Posts - Add some random posts with fuzz text to current channel. /test posts [fuzz] @@ -135,6 +138,10 @@ func (me *LoadTestProvider) DoCommand(a *App, args *model.CommandArgs, message s return me.PostCommand(a, args, message) } + if strings.HasPrefix(message, "threaded_post") { + return me.ThreadedPostCommand(a, args, message) + } + if strings.HasPrefix(message, "url") { return me.UrlCommand(a, args, message) } @@ -301,6 +308,34 @@ func (me *LoadTestProvider) ChannelsCommand(a *App, args *model.CommandArgs, mes return &model.CommandResponse{Text: "Added channels", ResponseType: model.COMMAND_RESPONSE_TYPE_EPHEMERAL} } +func (me *LoadTestProvider) ThreadedPostCommand(a *App, args *model.CommandArgs, message string) *model.CommandResponse { + var usernames []string + options := &model.UserGetOptions{InTeamId: args.TeamId, Page: 0, PerPage: 1000} + if profileUsers, err := a.Srv.Store.User().GetProfiles(options); err == nil { + usernames = make([]string, len(profileUsers)) + i := 0 + for _, userprof := range profileUsers { + usernames[i] = userprof.Username + i++ + } + } + + client := model.NewAPIv4Client(args.SiteURL) + client.MockSession(args.Session.Token) + testPoster := NewAutoPostCreator(client, args.ChannelId) + testPoster.Fuzzy = true + testPoster.Users = usernames + rpost, ok := testPoster.CreateRandomPost() + if !ok { + return &model.CommandResponse{Text: "Cannot create a post", ResponseType: model.COMMAND_RESPONSE_TYPE_EPHEMERAL} + } + for i := 0; i < 1000; i++ { + testPoster.CreateRandomPostNested(rpost.Id, rpost.Id) + } + + return &model.CommandResponse{Text: "Added threaded post", ResponseType: model.COMMAND_RESPONSE_TYPE_EPHEMERAL} +} + func (me *LoadTestProvider) PostsCommand(a *App, args *model.CommandArgs, message string) *model.CommandResponse { cmd := strings.TrimSpace(strings.TrimPrefix(message, "posts")) diff --git a/app/file.go b/app/file.go index bd2fb59aea..3dc20d4f37 100644 --- a/app/file.go +++ b/app/file.go @@ -290,7 +290,7 @@ func (a *App) MigrateFilenamesToFileInfos(post *model.Post) []*model.FileInfo { fileMigrationLock.Lock() defer fileMigrationLock.Unlock() - result, err := a.Srv.Store.Post().Get(post.Id) + result, err := a.Srv.Store.Post().Get(post.Id, false) if err != nil { mlog.Error("Unable to get post when migrating post to use FileInfos", mlog.Err(err), mlog.String("post_id", post.Id)) return []*model.FileInfo{} diff --git a/app/plugin_api.go b/app/plugin_api.go index 8c79ddf1d7..fa99755c62 100644 --- a/app/plugin_api.go +++ b/app/plugin_api.go @@ -501,7 +501,7 @@ func (api *PluginAPI) DeletePost(postId string) *model.AppError { } func (api *PluginAPI) GetPostThread(postId string) (*model.PostList, *model.AppError) { - return api.app.GetPostThread(postId) + return api.app.GetPostThread(postId, false) } func (api *PluginAPI) GetPost(postId string) (*model.Post, *model.AppError) { @@ -509,19 +509,19 @@ func (api *PluginAPI) GetPost(postId string) (*model.Post, *model.AppError) { } func (api *PluginAPI) GetPostsSince(channelId string, time int64) (*model.PostList, *model.AppError) { - return api.app.GetPostsSince(channelId, time) + return api.app.GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: time}) } func (api *PluginAPI) GetPostsAfter(channelId, postId string, page, perPage int) (*model.PostList, *model.AppError) { - return api.app.GetPostsAfterPost(channelId, postId, page, perPage) + return api.app.GetPostsAfterPost(model.GetPostsOptions{ChannelId: channelId, PostId: postId, Page: page, PerPage: perPage}) } func (api *PluginAPI) GetPostsBefore(channelId, postId string, page, perPage int) (*model.PostList, *model.AppError) { - return api.app.GetPostsBeforePost(channelId, postId, page, perPage) + return api.app.GetPostsBeforePost(model.GetPostsOptions{ChannelId: channelId, PostId: postId, Page: page, PerPage: perPage}) } func (api *PluginAPI) GetPostsForChannel(channelId string, page, perPage int) (*model.PostList, *model.AppError) { - return api.app.GetPostsPage(channelId, page, perPage) + return api.app.GetPostsPage(model.GetPostsOptions{ChannelId: channelId, Page: perPage, PerPage: page}) } func (api *PluginAPI) UpdatePost(post *model.Post) (*model.Post, *model.AppError) { diff --git a/app/post.go b/app/post.go index e226399136..07ed78b461 100644 --- a/app/post.go +++ b/app/post.go @@ -167,7 +167,7 @@ func (a *App) CreatePost(post *model.Post, channel *model.Channel, triggerWebhoo if len(post.RootId) > 0 { pchan = make(chan store.StoreResult, 1) go func() { - r, pErr := a.Srv.Store.Post().Get(post.RootId) + r, pErr := a.Srv.Store.Post().Get(post.RootId, false) pchan <- store.StoreResult{Data: r, Err: pErr} close(pchan) }() @@ -475,7 +475,7 @@ func (a *App) DeleteEphemeralPost(userId, postId string) { func (a *App) UpdatePost(post *model.Post, safeUpdate bool) (*model.Post, *model.AppError) { post.SanitizeProps() - postLists, err := a.Srv.Store.Post().Get(post.Id) + postLists, err := a.Srv.Store.Post().Get(post.Id, false) if err != nil { return nil, err } @@ -614,28 +614,28 @@ func (a *App) PatchPost(postId string, patch *model.PostPatch) (*model.Post, *mo return updatedPost, nil } -func (a *App) GetPostsPage(channelId string, page int, perPage int) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().GetPosts(channelId, page*perPage, perPage, true) +func (a *App) GetPostsPage(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + return a.Srv.Store.Post().GetPosts(options, false) } func (a *App) GetPosts(channelId string, offset int, limit int) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().GetPosts(channelId, offset, limit, true) + return a.Srv.Store.Post().GetPosts(model.GetPostsOptions{ChannelId: channelId, Page: offset, PerPage: limit}, true) } func (a *App) GetPostsEtag(channelId string) string { return a.Srv.Store.Post().GetEtag(channelId, true) } -func (a *App) GetPostsSince(channelId string, time int64) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().GetPostsSince(channelId, time, true) +func (a *App) GetPostsSince(options model.GetPostsSinceOptions) (*model.PostList, *model.AppError) { + return a.Srv.Store.Post().GetPostsSince(options, true) } func (a *App) GetSinglePost(postId string) (*model.Post, *model.AppError) { return a.Srv.Store.Post().GetSingle(postId) } -func (a *App) GetPostThread(postId string) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().Get(postId) +func (a *App) GetPostThread(postId string, skipFetchThreads bool) (*model.PostList, *model.AppError) { + return a.Srv.Store.Post().Get(postId, skipFetchThreads) } func (a *App) GetFlaggedPosts(userId string, offset int, limit int) (*model.PostList, *model.AppError) { @@ -651,7 +651,7 @@ func (a *App) GetFlaggedPostsForChannel(userId, channelId string, offset int, li } func (a *App) GetPermalinkPost(postId string, userId string) (*model.PostList, *model.AppError) { - list, err := a.Srv.Store.Post().Get(postId) + list, err := a.Srv.Store.Post().Get(postId, false) if err != nil { return nil, err } @@ -673,19 +673,19 @@ func (a *App) GetPermalinkPost(postId string, userId string) (*model.PostList, * return list, nil } -func (a *App) GetPostsBeforePost(channelId, postId string, page, perPage int) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().GetPostsBefore(channelId, postId, perPage, page*perPage) +func (a *App) GetPostsBeforePost(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + return a.Srv.Store.Post().GetPostsBefore(options) } -func (a *App) GetPostsAfterPost(channelId, postId string, page, perPage int) (*model.PostList, *model.AppError) { - return a.Srv.Store.Post().GetPostsAfter(channelId, postId, perPage, page*perPage) +func (a *App) GetPostsAfterPost(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + return a.Srv.Store.Post().GetPostsAfter(options) } -func (a *App) GetPostsAroundPost(postId, channelId string, offset, limit int, before bool) (*model.PostList, *model.AppError) { +func (a *App) GetPostsAroundPost(before bool, options model.GetPostsOptions) (*model.PostList, *model.AppError) { if before { - return a.Srv.Store.Post().GetPostsBefore(channelId, postId, limit, offset) + return a.Srv.Store.Post().GetPostsBefore(options) } - return a.Srv.Store.Post().GetPostsAfter(channelId, postId, limit, offset) + return a.Srv.Store.Post().GetPostsAfter(options) } func (a *App) GetPostAfterTime(channelId string, time int64) (*model.Post, *model.AppError) { @@ -773,8 +773,7 @@ func (a *App) AddCursorIdsForPostList(originalList *model.PostList, afterPost, b originalList.NextPostId = nextPostId originalList.PrevPostId = prevPostId } - -func (a *App) GetPostsForChannelAroundLastUnread(channelId, userId string, limitBefore, limitAfter int) (*model.PostList, *model.AppError) { +func (a *App) GetPostsForChannelAroundLastUnread(channelId, userId string, limitBefore, limitAfter int, skipFetchThreads bool) (*model.PostList, *model.AppError) { var member *model.ChannelMember var err *model.AppError if member, err = a.GetChannelMember(channelId, userId); err != nil { @@ -790,7 +789,7 @@ func (a *App) GetPostsForChannelAroundLastUnread(channelId, userId string, limit return model.NewPostList(), nil } - postList, err := a.GetPostThread(lastUnreadPostId) + postList, err := a.GetPostThread(lastUnreadPostId, skipFetchThreads) if err != nil { return nil, err } @@ -798,13 +797,13 @@ func (a *App) GetPostsForChannelAroundLastUnread(channelId, userId string, limit // channel organically, those replies will be added below. postList.Order = []string{lastUnreadPostId} - if postListBefore, err := a.GetPostsBeforePost(channelId, lastUnreadPostId, PAGE_DEFAULT, limitBefore); err != nil { + if postListBefore, err := a.GetPostsBeforePost(model.GetPostsOptions{ChannelId: channelId, PostId: lastUnreadPostId, Page: PAGE_DEFAULT, PerPage: limitBefore, SkipFetchThreads: skipFetchThreads}); err != nil { return nil, err } else if postListBefore != nil { postList.Extend(postListBefore) } - if postListAfter, err := a.GetPostsAfterPost(channelId, lastUnreadPostId, PAGE_DEFAULT, limitAfter-1); err != nil { + if postListAfter, err := a.GetPostsAfterPost(model.GetPostsOptions{ChannelId: channelId, PostId: lastUnreadPostId, Page: PAGE_DEFAULT, PerPage: limitAfter - 1, SkipFetchThreads: skipFetchThreads}); err != nil { return nil, err } else if postListAfter != nil { postList.Extend(postListAfter) @@ -1216,7 +1215,7 @@ func (a *App) countMentionsFromPost(user *model.User, post *model.Post) (int, *m // A mapping of thread root IDs to whether or not a post in that thread mentions the user mentionedByThread := make(map[string]bool) - thread, err := a.GetPostThread(post.Id) + thread, err := a.GetPostThread(post.Id, false) if err != nil { return 0, err } @@ -1230,7 +1229,12 @@ func (a *App) countMentionsFromPost(user *model.User, post *model.Post) (int, *m page := 0 perPage := 200 for { - postList, err := a.GetPostsAfterPost(post.ChannelId, post.Id, page, perPage) + postList, err := a.GetPostsAfterPost(model.GetPostsOptions{ + ChannelId: post.ChannelId, + PostId: post.Id, + Page: page, + PerPage: perPage, + }) if err != nil { return 0, err } diff --git a/model/post.go b/model/post.go index b424b3f419..06899d20f5 100644 --- a/model/post.go +++ b/model/post.go @@ -74,7 +74,6 @@ type Post struct { OriginalId string `json:"original_id"` Message string `json:"message"` - // MessageSource will contain the message as submitted by the user if Message has been modified // by Mattermost for presentation (e.g if an image proxy is being used). It should be used to // populate edit boxes if present. @@ -89,7 +88,8 @@ type Post struct { HasReactions bool `json:"has_reactions,omitempty"` // Transient data populated before sending a post to the client - Metadata *PostMetadata `json:"metadata,omitempty" db:"-"` + ReplyCount int64 `json:"reply_count" db:"-"` + Metadata *PostMetadata `json:"metadata,omitempty" db:"-"` } type PostEphemeral struct { @@ -171,6 +171,20 @@ func (o *Post) ToUnsanitizedJson() string { return string(b) } +type GetPostsSinceOptions struct { + ChannelId string + Time int64 + SkipFetchThreads bool +} + +type GetPostsOptions struct { + ChannelId string + PostId string + Page int + PerPage int + SkipFetchThreads bool +} + func PostFromJson(data io.Reader) *Post { var o *Post json.NewDecoder(data).Decode(&o) diff --git a/store/localcachelayer/main_test.go b/store/localcachelayer/main_test.go index 35531668ff..2f75014c4d 100644 --- a/store/localcachelayer/main_test.go +++ b/store/localcachelayer/main_test.go @@ -181,18 +181,25 @@ func getMockStore() *mocks.Store { mockChannelStore.On("GetPinnedPostCount", "id", false).Return(mockPinnedPostsCount, nil) fakePosts := &model.PostList{} + fakeOptions := model.GetPostsOptions{ChannelId: "123", PerPage: 30} mockPostStore := mocks.PostStore{} - mockPostStore.On("GetPosts", "123", 0, 30, true).Return(fakePosts, nil) - mockPostStore.On("GetPosts", "123", 0, 30, false).Return(fakePosts, nil) + mockPostStore.On("GetPosts", fakeOptions, true).Return(fakePosts, nil) + mockPostStore.On("GetPosts", fakeOptions, false).Return(fakePosts, nil) mockPostStore.On("InvalidateLastPostTimeCache", "12360") + mockPostStoreOptions := model.GetPostsSinceOptions{ + ChannelId: "channelId", + Time: 1, + SkipFetchThreads: false, + } + mockPostStoreEtagResult := fmt.Sprintf("%v.%v", model.CurrentVersion, 1) mockPostStore.On("ClearCaches") mockPostStore.On("InvalidateLastPostTimeCache", "channelId") mockPostStore.On("GetEtag", "channelId", true).Return(mockPostStoreEtagResult) mockPostStore.On("GetEtag", "channelId", false).Return(mockPostStoreEtagResult) - mockPostStore.On("GetPostsSince", "channelId", int64(1), true).Return(model.NewPostList(), nil) - mockPostStore.On("GetPostsSince", "channelId", int64(1), false).Return(model.NewPostList(), nil) + mockPostStore.On("GetPostsSince", mockPostStoreOptions, true).Return(model.NewPostList(), nil) + mockPostStore.On("GetPostsSince", mockPostStoreOptions, false).Return(model.NewPostList(), nil) mockStore.On("Post").Return(&mockPostStore) fakeTermsOfService := model.TermsOfService{Id: "123", CreateAt: 11111, UserId: "321", Text: "Terms of service test"} diff --git a/store/localcachelayer/post_layer.go b/store/localcachelayer/post_layer.go index db12ca2d84..c5219f48d9 100644 --- a/store/localcachelayer/post_layer.go +++ b/store/localcachelayer/post_layer.go @@ -78,51 +78,52 @@ func (s LocalCachePostStore) GetEtag(channelId string, allowFromCache bool) stri return result } -func (s LocalCachePostStore) GetPostsSince(channelId string, time int64, allowFromCache bool) (*model.PostList, *model.AppError) { +func (s LocalCachePostStore) GetPostsSince(options model.GetPostsSinceOptions, allowFromCache bool) (*model.PostList, *model.AppError) { if allowFromCache { // If the last post in the channel's time is less than or equal to the time we are getting posts since, // we can safely return no posts. - if lastTime := s.rootStore.doStandardReadCache(s.rootStore.lastPostTimeCache, channelId); lastTime != nil && lastTime.(int64) <= time { + if lastTime := s.rootStore.doStandardReadCache(s.rootStore.lastPostTimeCache, options.ChannelId); lastTime != nil && lastTime.(int64) <= options.Time { list := model.NewPostList() return list, nil } } - list, err := s.PostStore.GetPostsSince(channelId, time, allowFromCache) + list, err := s.PostStore.GetPostsSince(options, allowFromCache) - latestUpdate := time + latestUpdate := options.Time if err == nil { for _, p := range list.ToSlice() { if latestUpdate < p.UpdateAt { latestUpdate = p.UpdateAt } } - s.rootStore.doStandardAddToCache(s.rootStore.lastPostTimeCache, channelId, latestUpdate) + s.rootStore.doStandardAddToCache(s.rootStore.lastPostTimeCache, options.ChannelId, latestUpdate) } return list, err } -func (s LocalCachePostStore) GetPosts(channelId string, offset int, limit int, allowFromCache bool) (*model.PostList, *model.AppError) { +func (s LocalCachePostStore) GetPosts(options model.GetPostsOptions, allowFromCache bool) (*model.PostList, *model.AppError) { if !allowFromCache { - return s.PostStore.GetPosts(channelId, offset, limit, allowFromCache) + return s.PostStore.GetPosts(options, allowFromCache) } + offset := options.PerPage * options.Page // Caching only occurs on limits of 30 and 60, the common limits requested by MM clients - if offset == 0 && (limit == 60 || limit == 30) { - if cacheItem := s.rootStore.doStandardReadCache(s.rootStore.postLastPostsCache, fmt.Sprintf("%s%v", channelId, limit)); cacheItem != nil { + if offset == 0 && (options.PerPage == 60 || options.PerPage == 30) { + if cacheItem := s.rootStore.doStandardReadCache(s.rootStore.postLastPostsCache, fmt.Sprintf("%s%v", options.ChannelId, options.PerPage)); cacheItem != nil { return cacheItem.(*model.PostList), nil } } - list, err := s.PostStore.GetPosts(channelId, offset, limit, allowFromCache) + list, err := s.PostStore.GetPosts(options, false) if err != nil { return nil, err } // Caching only occurs on limits of 30 and 60, the common limits requested by MM clients - if offset == 0 && (limit == 60 || limit == 30) { - s.rootStore.doStandardAddToCache(s.rootStore.postLastPostsCache, fmt.Sprintf("%s%v", channelId, limit), list) + if offset == 0 && (options.PerPage == 60 || options.PerPage == 30) { + s.rootStore.doStandardAddToCache(s.rootStore.postLastPostsCache, fmt.Sprintf("%s%v", options.ChannelId, options.PerPage), list) } return list, err diff --git a/store/localcachelayer/post_layer_test.go b/store/localcachelayer/post_layer_test.go index c872648582..7d832be621 100644 --- a/store/localcachelayer/post_layer_test.go +++ b/store/localcachelayer/post_layer_test.go @@ -21,6 +21,11 @@ func TestPostStore(t *testing.T) { func TestPostStoreLastPostTimeCache(t *testing.T) { var fakeLastTime int64 = 1 channelId := "channelId" + fakeOptions := model.GetPostsSinceOptions{ + ChannelId: channelId, + Time: fakeLastTime, + SkipFetchThreads: false, + } t.Run("GetEtag: first call not cached, second cached and returning same data", func(t *testing.T) { mockStore := getMockStore() @@ -80,12 +85,12 @@ func TestPostStoreLastPostTimeCache(t *testing.T) { expectedResult := model.NewPostList() - list, err := cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + list, err := cachedStore.Post().GetPostsSince(fakeOptions, true) require.Nil(t, err) assert.Equal(t, list, expectedResult) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 1) - list, err = cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + list, err = cachedStore.Post().GetPostsSince(fakeOptions, true) require.Nil(t, err) assert.Equal(t, list, expectedResult) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 1) @@ -96,9 +101,9 @@ func TestPostStoreLastPostTimeCache(t *testing.T) { mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + cachedStore.Post().GetPostsSince(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 1) - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, false) + cachedStore.Post().GetPostsSince(fakeOptions, false) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 2) }) @@ -107,10 +112,10 @@ func TestPostStoreLastPostTimeCache(t *testing.T) { mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + cachedStore.Post().GetPostsSince(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 1) cachedStore.Post().InvalidateLastPostTimeCache(channelId) - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + cachedStore.Post().GetPostsSince(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 2) }) @@ -119,28 +124,29 @@ func TestPostStoreLastPostTimeCache(t *testing.T) { mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + cachedStore.Post().GetPostsSince(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 1) cachedStore.Post().ClearCaches() - cachedStore.Post().GetPostsSince(channelId, fakeLastTime, true) + cachedStore.Post().GetPostsSince(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPostsSince", 2) }) } func TestPostStoreCache(t *testing.T) { fakePosts := &model.PostList{} + fakeOptions := model.GetPostsOptions{ChannelId: "123", PerPage: 30} t.Run("first call not cached, second cached and returning same data", func(t *testing.T) { mockStore := getMockStore() mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - gotPosts, err := cachedStore.Post().GetPosts("123", 0, 30, true) + gotPosts, err := cachedStore.Post().GetPosts(fakeOptions, true) require.Nil(t, err) assert.Equal(t, fakePosts, gotPosts) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 1) - _, _ = cachedStore.Post().GetPosts("123", 0, 30, true) + _, _ = cachedStore.Post().GetPosts(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 1) }) @@ -149,12 +155,12 @@ func TestPostStoreCache(t *testing.T) { mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - gotPosts, err := cachedStore.Post().GetPosts("123", 0, 30, true) + gotPosts, err := cachedStore.Post().GetPosts(fakeOptions, true) require.Nil(t, err) assert.Equal(t, fakePosts, gotPosts) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 1) - _, _ = cachedStore.Post().GetPosts("123", 0, 30, false) + _, _ = cachedStore.Post().GetPosts(fakeOptions, false) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 2) }) @@ -163,14 +169,14 @@ func TestPostStoreCache(t *testing.T) { mockCacheProvider := getMockCacheProvider() cachedStore := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider) - gotPosts, err := cachedStore.Post().GetPosts("123", 0, 30, true) + gotPosts, err := cachedStore.Post().GetPosts(fakeOptions, true) require.Nil(t, err) assert.Equal(t, fakePosts, gotPosts) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 1) cachedStore.Post().InvalidateLastPostTimeCache("12360") - _, _ = cachedStore.Post().GetPosts("123", 0, 30, true) + _, _ = cachedStore.Post().GetPosts(fakeOptions, true) mockStore.Post().(*mocks.PostStore).AssertNumberOfCalls(t, "GetPosts", 1) }) diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index e3324819c0..e9d8bdf424 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -105,6 +105,12 @@ func (s *SqlPostStore) Save(post *model.Post) (*model.Post, *model.AppError) { if _, err := s.GetMaster().Exec("UPDATE Posts SET UpdateAt = :UpdateAt WHERE Id = :RootId", map[string]interface{}{"UpdateAt": time, "RootId": post.RootId}); err != nil { mlog.Error("Error updating Post UpdateAt.", mlog.Err(err)) } + } else { + if count, err := s.GetMaster().SelectInt("SELECT COUNT(*) FROM Posts WHERE RootId = :Id", map[string]interface{}{"Id": post.Id}); err != nil { + mlog.Error("Error fetching post's thread.", mlog.Err(err)) + } else { + post.ReplyCount = count + } } return post, nil @@ -162,7 +168,7 @@ func (s *SqlPostStore) GetFlaggedPosts(userId string, offset int, limit int) (*m pl := model.NewPostList() var posts []*model.Post - if _, err := s.GetReplica().Select(&posts, "SELECT * FROM Posts WHERE Id IN (SELECT Name FROM Preferences WHERE UserId = :UserId AND Category = :Category) AND DeleteAt = 0 ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset", map[string]interface{}{"UserId": userId, "Category": model.PREFERENCE_CATEGORY_FLAGGED_POST, "Offset": offset, "Limit": limit}); err != nil { + if _, err := s.GetReplica().Select(&posts, "SELECT *, (SELECT count(Posts.Id) FROM Posts WHERE Posts.RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount FROM Posts p WHERE Id IN (SELECT Name FROM Preferences WHERE UserId = :UserId AND Category = :Category) AND DeleteAt = 0 ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset", map[string]interface{}{"UserId": userId, "Category": model.PREFERENCE_CATEGORY_FLAGGED_POST, "Offset": offset, "Limit": limit}); err != nil { return nil, model.NewAppError("SqlPostStore.GetFlaggedPosts", "store.sql_post.get_flagged_posts.app_error", nil, err.Error(), http.StatusInternalServerError) } @@ -181,7 +187,7 @@ func (s *SqlPostStore) GetFlaggedPostsForTeam(userId, teamId string, offset int, query := ` SELECT - A.* + A.*, (SELECT count(Posts.Id) FROM Posts WHERE Posts.RootId = A.Id AND Posts.DeleteAt = 0) as ReplyCount FROM (SELECT * @@ -223,8 +229,8 @@ func (s *SqlPostStore) GetFlaggedPostsForChannel(userId, channelId string, offse var posts []*model.Post query := ` SELECT - * - FROM Posts + *, (SELECT count(Posts.Id) FROM Posts WHERE Posts.RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount + FROM Posts p WHERE Id IN (SELECT Name FROM Preferences WHERE UserId = :UserId AND Category = :Category) AND ChannelId = :ChannelId @@ -243,7 +249,7 @@ func (s *SqlPostStore) GetFlaggedPostsForChannel(userId, channelId string, offse return pl, nil } -func (s *SqlPostStore) Get(id string) (*model.PostList, *model.AppError) { +func (s *SqlPostStore) Get(id string, skipFetchThreads bool) (*model.PostList, *model.AppError) { pl := model.NewPostList() if len(id) == 0 { @@ -251,35 +257,35 @@ func (s *SqlPostStore) Get(id string) (*model.PostList, *model.AppError) { } var post model.Post - err := s.GetReplica().SelectOne(&post, "SELECT * FROM Posts WHERE Id = :Id AND DeleteAt = 0", map[string]interface{}{"Id": id}) + postFetchQuery := "SELECT p.*, (SELECT count(Posts.Id) FROM Posts WHERE Posts.RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount FROM Posts p WHERE p.Id = :Id AND p.DeleteAt = 0" + err := s.GetReplica().SelectOne(&post, postFetchQuery, map[string]interface{}{"Id": id}) if err != nil { return nil, model.NewAppError("SqlPostStore.GetPost", "store.sql_post.get.app_error", nil, "id="+id+err.Error(), http.StatusNotFound) } - pl.AddPost(&post) pl.AddOrder(id) + if !skipFetchThreads { + rootId := post.RootId - rootId := post.RootId + if rootId == "" { + rootId = post.Id + } - if rootId == "" { - rootId = post.Id + if len(rootId) == 0 { + return nil, model.NewAppError("SqlPostStore.GetPost", "store.sql_post.get.app_error", nil, "root_id="+rootId, http.StatusInternalServerError) + } + + var posts []*model.Post + _, err = s.GetReplica().Select(&posts, "SELECT *, (SELECT count(Id) FROM Posts WHERE RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount FROM Posts p WHERE (Id = :Id OR RootId = :RootId) AND DeleteAt = 0", map[string]interface{}{"Id": rootId, "RootId": rootId}) + if err != nil { + return nil, model.NewAppError("SqlPostStore.GetPost", "store.sql_post.get.app_error", nil, "root_id="+rootId+err.Error(), http.StatusInternalServerError) + } + + for _, p := range posts { + pl.AddPost(p) + pl.AddOrder(p.Id) + } } - - if len(rootId) == 0 { - return nil, model.NewAppError("SqlPostStore.GetPost", "store.sql_post.get.app_error", nil, "root_id="+rootId, http.StatusInternalServerError) - } - - var posts []*model.Post - _, err = s.GetReplica().Select(&posts, "SELECT * FROM Posts WHERE (Id = :Id OR RootId = :RootId) AND DeleteAt = 0", map[string]interface{}{"Id": rootId, "RootId": rootId}) - if err != nil { - return nil, model.NewAppError("SqlPostStore.GetPost", "store.sql_post.get.app_error", nil, "root_id="+rootId+err.Error(), http.StatusInternalServerError) - } - - for _, p := range posts { - pl.AddPost(p) - pl.AddOrder(p.Id) - } - return pl, nil } @@ -394,20 +400,21 @@ func (s *SqlPostStore) PermanentDeleteByChannel(channelId string) *model.AppErro return nil } -func (s *SqlPostStore) GetPosts(channelId string, offset int, limit int, allowFromCache bool) (*model.PostList, *model.AppError) { - if limit > 1000 { - return nil, model.NewAppError("SqlPostStore.GetLinearPosts", "store.sql_post.get_posts.app_error", nil, "channelId="+channelId, http.StatusBadRequest) +func (s *SqlPostStore) GetPosts(options model.GetPostsOptions, _ bool) (*model.PostList, *model.AppError) { + if options.PerPage > 1000 { + return nil, model.NewAppError("SqlPostStore.GetLinearPosts", "store.sql_post.get_posts.app_error", nil, "channelId="+options.ChannelId, http.StatusBadRequest) } + offset := options.PerPage * options.Page rpc := make(chan store.StoreResult, 1) go func() { - posts, err := s.getRootPosts(channelId, offset, limit) + posts, err := s.getRootPosts(options.ChannelId, offset, options.PerPage, options.SkipFetchThreads) rpc <- store.StoreResult{Data: posts, Err: err} close(rpc) }() cpc := make(chan store.StoreResult, 1) go func() { - posts, err := s.getParentsPosts(channelId, offset, limit) + posts, err := s.getParentsPosts(options.ChannelId, offset, options.PerPage, options.SkipFetchThreads) cpc <- store.StoreResult{Data: posts, Err: err} close(cpc) }() @@ -442,15 +449,20 @@ func (s *SqlPostStore) GetPosts(channelId string, offset int, limit int, allowFr return list, err } -func (s *SqlPostStore) GetPostsSince(channelId string, time int64, allowFromCache bool) (*model.PostList, *model.AppError) { - if s.metrics != nil { - s.metrics.IncrementMemCacheMissCounter("Last Post Time") +func (s *SqlPostStore) GetPostsSince(options model.GetPostsSinceOptions, allowFromCache bool) (*model.PostList, *model.AppError) { + var posts []*model.Post + + replyCountQuery1 := "" + replyCountQuery2 := "" + if options.SkipFetchThreads { + replyCountQuery1 = `, (SELECT COUNT(Posts.Id) FROM Posts WHERE p1.RootId = '' AND Posts.RootId = p1.Id AND Posts.DeleteAt = 0) as ReplyCount` + replyCountQuery2 = `, (SELECT COUNT(Posts.Id) FROM Posts WHERE p2.RootId = '' AND Posts.RootId = p2.Id AND Posts.DeleteAt = 0) as ReplyCount` } var query string - var posts []*model.Post + // union of IDs and then join to get full posts is faster in mysql if s.DriverName() == model.DATABASE_DRIVER_MYSQL { - query = `SELECT * FROM Posts p1 JOIN ( + query = `SELECT *` + replyCountQuery1 + ` FROM Posts p1 JOIN ( (SELECT Id FROM @@ -480,7 +492,7 @@ func (s *SqlPostStore) GetPostsSince(channelId string, time int64, allowFromCach } else if s.DriverName() == model.DATABASE_DRIVER_POSTGRES { query = ` (SELECT - * + *` + replyCountQuery1 + ` FROM Posts p1 WHERE @@ -489,7 +501,7 @@ func (s *SqlPostStore) GetPostsSince(channelId string, time int64, allowFromCach LIMIT 1000) UNION (SELECT - * + *` + replyCountQuery2 + ` FROM Posts p2 WHERE @@ -505,17 +517,17 @@ func (s *SqlPostStore) GetPostsSince(channelId string, time int64, allowFromCach LIMIT 1000) temp_tab)) ORDER BY CreateAt DESC` } - _, err := s.GetReplica().Select(&posts, query, map[string]interface{}{"ChannelId": channelId, "Time": time}) + _, err := s.GetReplica().Select(&posts, query, map[string]interface{}{"ChannelId": options.ChannelId, "Time": options.Time}) if err != nil { - return nil, model.NewAppError("SqlPostStore.GetPostsSince", "store.sql_post.get_posts_since.app_error", nil, "channelId="+channelId+err.Error(), http.StatusInternalServerError) + return nil, model.NewAppError("SqlPostStore.GetPostsSince", "store.sql_post.get_posts_since.app_error", nil, "channelId="+options.ChannelId+err.Error(), http.StatusInternalServerError) } list := model.NewPostList() for _, p := range posts { list.AddPost(p) - if p.UpdateAt > time { + if p.UpdateAt > options.Time { list.AddOrder(p.Id) } } @@ -523,16 +535,20 @@ func (s *SqlPostStore) GetPostsSince(channelId string, time int64, allowFromCach return list, nil } -func (s *SqlPostStore) GetPostsBefore(channelId string, postId string, limit int, offset int) (*model.PostList, *model.AppError) { - return s.getPostsAround(channelId, postId, limit, offset, true) +func (s *SqlPostStore) GetPostsBefore(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + return s.getPostsAround(true, options) } -func (s *SqlPostStore) GetPostsAfter(channelId string, postId string, limit int, offset int) (*model.PostList, *model.AppError) { - return s.getPostsAround(channelId, postId, limit, offset, false) +func (s *SqlPostStore) GetPostsAfter(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + return s.getPostsAround(false, options) } -func (s *SqlPostStore) getPostsAround(channelId string, postId string, limit int, offset int, before bool) (*model.PostList, *model.AppError) { - var direction, sort string +func (s *SqlPostStore) getPostsAround(before bool, options model.GetPostsOptions) (*model.PostList, *model.AppError) { + offset := options.Page * options.PerPage + var posts, parents []*model.Post + + var direction string + var sort string if before { direction = "<" sort = "DESC" @@ -540,23 +556,29 @@ func (s *SqlPostStore) getPostsAround(channelId string, postId string, limit int direction = ">" sort = "ASC" } + replyCountSubQuery := s.getQueryBuilder().Select("COUNT(Posts.Id)").From("Posts").Where(sq.Expr("p.RootId = '' AND RootId = p.Id AND DeleteAt = 0")) + query := s.getQueryBuilder().Select("p.*") + if options.SkipFetchThreads { + query = query.Column(sq.Alias(replyCountSubQuery, "ReplyCount")) + } + query = query.From("Posts p"). + Where(sq.And{ + sq.Expr(`CreateAt `+direction+` (SELECT CreateAt FROM Posts WHERE Id = ?)`, options.PostId), + sq.Eq{"ChannelId": options.ChannelId}, + sq.Eq{"DeleteAt": int(0)}, + }). + OrderBy("CreateAt " + sort). + Limit(uint64(options.PerPage)). + Offset(uint64(offset)) + + queryString, args, err := query.ToSql() - var posts, parents []*model.Post - _, err := s.GetReplica().Select(&posts, - `SELECT - * - FROM - Posts - WHERE - CreateAt `+direction+` (SELECT CreateAt FROM Posts WHERE Id = :PostId) - AND ChannelId = :ChannelId - AND DeleteAt = 0 - ORDER BY CreateAt `+sort+` - LIMIT :Limit - OFFSET :Offset`, - map[string]interface{}{"ChannelId": channelId, "PostId": postId, "Limit": limit, "Offset": offset}) if err != nil { - return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get.app_error", nil, "channelId="+channelId+err.Error(), http.StatusInternalServerError) + return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get.app_error", nil, "channelId="+options.ChannelId+err.Error(), http.StatusInternalServerError) + } + _, err = s.GetMaster().Select(&posts, queryString, args...) + if err != nil { + return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get.app_error", nil, "channelId="+options.ChannelId+err.Error(), http.StatusInternalServerError) } if len(posts) > 0 { @@ -567,28 +589,32 @@ func (s *SqlPostStore) getPostsAround(channelId string, postId string, limit int rootIds = append(rootIds, post.RootId) } } + rootQuery := s.getQueryBuilder().Select("p.*") + idQuery := sq.Or{ + sq.Eq{"Id": rootIds}, + } + if options.SkipFetchThreads { + rootQuery = rootQuery.Column(sq.Alias(replyCountSubQuery, "ReplyCount")) + } else { + idQuery = append(idQuery, sq.Eq{"RootId": rootIds}) // preserve original behaviour + } - keys, params := MapStringsToQueryParams(rootIds, "PostId") + rootQuery = rootQuery.From("Posts p"). + Where(sq.And{ + idQuery, + sq.Eq{"ChannelId": options.ChannelId}, + sq.Eq{"DeleteAt": 0}, + }). + OrderBy("CreateAt DESC") - params["ChannelId"] = channelId - params["PostId"] = postId - params["Limit"] = limit - params["Offset"] = offset - - _, err = s.GetReplica().Select(&parents, - `SELECT - * - FROM - Posts - WHERE - (Id IN `+keys+` OR RootId IN `+keys+`) - AND ChannelId = :ChannelId - AND DeleteAt = 0 - ORDER BY CreateAt DESC`, - params) + rootQueryString, rootArgs, err := rootQuery.ToSql() if err != nil { - return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get_parent.app_error", nil, "channelId="+channelId+err.Error(), http.StatusInternalServerError) + return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get_parent.app_error", nil, "channelId="+options.ChannelId+err.Error(), http.StatusInternalServerError) + } + _, err = s.GetMaster().Select(&parents, rootQueryString, rootArgs...) + if err != nil { + return nil, model.NewAppError("SqlPostStore.GetPostContext", "store.sql_post.get_posts_around.get_parent.app_error", nil, "channelId="+options.ChannelId+err.Error(), http.StatusInternalServerError) } } @@ -687,18 +713,24 @@ func (s *SqlPostStore) GetPostAfterTime(channelId string, time int64) (*model.Po return post, nil } -func (s *SqlPostStore) getRootPosts(channelId string, offset int, limit int) ([]*model.Post, *model.AppError) { +func (s *SqlPostStore) getRootPosts(channelId string, offset int, limit int, skipFetchThreads bool) ([]*model.Post, *model.AppError) { var posts []*model.Post - _, err := s.GetReplica().Select(&posts, "SELECT * FROM Posts WHERE ChannelId = :ChannelId AND DeleteAt = 0 ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset", map[string]interface{}{"ChannelId": channelId, "Offset": offset, "Limit": limit}) + var fetchQuery string + if skipFetchThreads { + fetchQuery = "SELECT p.*, (SELECT COUNT(Posts.Id) FROM Posts WHERE p.RootId = '' AND Posts.RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount FROM Posts p WHERE ChannelId = :ChannelId AND DeleteAt = 0 ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset" + } else { + fetchQuery = "SELECT * FROM Posts WHERE ChannelId = :ChannelId AND DeleteAt = 0 ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset" + } + _, err := s.GetReplica().Select(&posts, fetchQuery, map[string]interface{}{"ChannelId": channelId, "Offset": offset, "Limit": limit}) if err != nil { return nil, model.NewAppError("SqlPostStore.GetLinearPosts", "store.sql_post.get_root_posts.app_error", nil, "channelId="+channelId+err.Error(), http.StatusInternalServerError) } return posts, nil } -func (s *SqlPostStore) getParentsPosts(channelId string, offset int, limit int) ([]*model.Post, *model.AppError) { +func (s *SqlPostStore) getParentsPosts(channelId string, offset int, limit int, skipFetchThreads bool) ([]*model.Post, *model.AppError) { if s.DriverName() == model.DATABASE_DRIVER_POSTGRES { - return s.getParentsPostsPostgreSQL(channelId, offset, limit) + return s.getParentsPostsPostgreSQL(channelId, offset, limit, skipFetchThreads) } // query parent Ids first @@ -736,10 +768,16 @@ func (s *SqlPostStore) getParentsPosts(channelId string, offset int, limit int) } placeholderString := strings.Join(placeholders, ", ") params["ChannelId"] = channelId - whereStatement := "p.Id IN (" + placeholderString + ") OR p.RootId IN (" + placeholderString + ")" + replyCountQuery := "" + whereStatement := "p.Id IN (" + placeholderString + ")" + if skipFetchThreads { + replyCountQuery = `, (SELECT COUNT(Posts.Id) FROM Posts WHERE p.RootId = '' AND Posts.RootId = p.Id AND Posts.DeleteAt = 0) as ReplyCount` + } else { + whereStatement += " OR p.RootId IN (" + placeholderString + ")" + } var posts []*model.Post _, err = s.GetReplica().Select(&posts, ` - SELECT p.* + SELECT p.*`+replyCountQuery+` FROM Posts p WHERE @@ -754,10 +792,17 @@ func (s *SqlPostStore) getParentsPosts(channelId string, offset int, limit int) return posts, nil } -func (s *SqlPostStore) getParentsPostsPostgreSQL(channelId string, offset int, limit int) ([]*model.Post, *model.AppError) { +func (s *SqlPostStore) getParentsPostsPostgreSQL(channelId string, offset int, limit int, skipFetchThreads bool) ([]*model.Post, *model.AppError) { var posts []*model.Post + replyCountQuery := "" + onStatement := "q1.RootId = q2.Id" + if skipFetchThreads { + replyCountQuery = ` ,(SELECT COUNT(Posts.Id) FROM Posts WHERE q2.RootId = '' AND Posts.RootId = q2.Id AND Posts.DeleteAt = 0) as ReplyCount` + } else { + onStatement += " OR q1.RootId = q2.RootId" + } _, err := s.GetReplica().Select(&posts, - `SELECT q2.* + `SELECT q2.*`+replyCountQuery+` FROM Posts q2 INNER JOIN @@ -774,7 +819,7 @@ func (s *SqlPostStore) getParentsPostsPostgreSQL(channelId string, offset int, l ORDER BY CreateAt DESC LIMIT :Limit OFFSET :Offset) q3 WHERE q3.RootId != '') q1 - ON q1.RootId = q2.Id OR q1.RootId = q2.RootId + ON `+onStatement+` WHERE ChannelId = :ChannelId2 AND DeleteAt = 0 @@ -942,9 +987,9 @@ func (s *SqlPostStore) Search(teamId string, userId string, params *model.Search searchQuery := ` SELECT - * + * ,(SELECT COUNT(Posts.Id) FROM Posts WHERE q2.RootId = '' AND Posts.RootId = q2.Id AND Posts.DeleteAt = 0) as ReplyCount FROM - Posts + Posts q2 WHERE DeleteAt = 0 AND Type NOT LIKE '` + model.POST_SYSTEM_MESSAGE_PREFIX + `%' diff --git a/store/store.go b/store/store.go index 4dbfce6ad6..4ab9a82219 100644 --- a/store/store.go +++ b/store/store.go @@ -210,18 +210,18 @@ type ChannelMemberHistoryStore interface { type PostStore interface { Save(post *model.Post) (*model.Post, *model.AppError) Update(newPost *model.Post, oldPost *model.Post) (*model.Post, *model.AppError) - Get(id string) (*model.PostList, *model.AppError) + Get(id string, skipFetchThreads bool) (*model.PostList, *model.AppError) GetSingle(id string) (*model.Post, *model.AppError) Delete(postId string, time int64, deleteByID string) *model.AppError PermanentDeleteByUser(userId string) *model.AppError PermanentDeleteByChannel(channelId string) *model.AppError - GetPosts(channelId string, offset int, limit int, allowFromCache bool) (*model.PostList, *model.AppError) + GetPosts(options model.GetPostsOptions, allowFromCache bool) (*model.PostList, *model.AppError) GetFlaggedPosts(userId string, offset int, limit int) (*model.PostList, *model.AppError) GetFlaggedPostsForTeam(userId, teamId string, offset int, limit int) (*model.PostList, *model.AppError) GetFlaggedPostsForChannel(userId, channelId string, offset int, limit int) (*model.PostList, *model.AppError) - GetPostsBefore(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) - GetPostsAfter(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) - GetPostsSince(channelId string, time int64, allowFromCache bool) (*model.PostList, *model.AppError) + GetPostsBefore(options model.GetPostsOptions) (*model.PostList, *model.AppError) + GetPostsAfter(options model.GetPostsOptions) (*model.PostList, *model.AppError) + GetPostsSince(options model.GetPostsSinceOptions, allowFromCache bool) (*model.PostList, *model.AppError) GetPostAfterTime(channelId string, time int64) (*model.Post, *model.AppError) GetPostIdAfterTime(channelId string, time int64) (string, *model.AppError) GetPostIdBeforeTime(channelId string, time int64) (string, *model.AppError) diff --git a/store/storetest/mocks/PostStore.go b/store/storetest/mocks/PostStore.go index ae38a5cb4d..bd879d60d0 100644 --- a/store/storetest/mocks/PostStore.go +++ b/store/storetest/mocks/PostStore.go @@ -108,13 +108,13 @@ func (_m *PostStore) Delete(postId string, time int64, deleteByID string) *model return r0 } -// Get provides a mock function with given fields: id -func (_m *PostStore) Get(id string) (*model.PostList, *model.AppError) { - ret := _m.Called(id) +// Get provides a mock function with given fields: id, skipFetchThreads +func (_m *PostStore) Get(id string, skipFetchThreads bool) (*model.PostList, *model.AppError) { + ret := _m.Called(id, skipFetchThreads) var r0 *model.PostList - if rf, ok := ret.Get(0).(func(string) *model.PostList); ok { - r0 = rf(id) + if rf, ok := ret.Get(0).(func(string, bool) *model.PostList); ok { + r0 = rf(id, skipFetchThreads) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostList) @@ -122,8 +122,8 @@ func (_m *PostStore) Get(id string) (*model.PostList, *model.AppError) { } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { - r1 = rf(id) + if rf, ok := ret.Get(1).(func(string, bool) *model.AppError); ok { + r1 = rf(id, skipFetchThreads) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) @@ -382,13 +382,13 @@ func (_m *PostStore) GetPostIdBeforeTime(channelId string, time int64) (string, return r0, r1 } -// GetPosts provides a mock function with given fields: channelId, offset, limit, allowFromCache -func (_m *PostStore) GetPosts(channelId string, offset int, limit int, allowFromCache bool) (*model.PostList, *model.AppError) { - ret := _m.Called(channelId, offset, limit, allowFromCache) +// GetPosts provides a mock function with given fields: options, allowFromCache +func (_m *PostStore) GetPosts(options model.GetPostsOptions, allowFromCache bool) (*model.PostList, *model.AppError) { + ret := _m.Called(options, allowFromCache) var r0 *model.PostList - if rf, ok := ret.Get(0).(func(string, int, int, bool) *model.PostList); ok { - r0 = rf(channelId, offset, limit, allowFromCache) + if rf, ok := ret.Get(0).(func(model.GetPostsOptions, bool) *model.PostList); ok { + r0 = rf(options, allowFromCache) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostList) @@ -396,8 +396,8 @@ func (_m *PostStore) GetPosts(channelId string, offset int, limit int, allowFrom } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string, int, int, bool) *model.AppError); ok { - r1 = rf(channelId, offset, limit, allowFromCache) + if rf, ok := ret.Get(1).(func(model.GetPostsOptions, bool) *model.AppError); ok { + r1 = rf(options, allowFromCache) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) @@ -407,13 +407,13 @@ func (_m *PostStore) GetPosts(channelId string, offset int, limit int, allowFrom return r0, r1 } -// GetPostsAfter provides a mock function with given fields: channelId, postId, numPosts, offset -func (_m *PostStore) GetPostsAfter(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) { - ret := _m.Called(channelId, postId, numPosts, offset) +// GetPostsAfter provides a mock function with given fields: options +func (_m *PostStore) GetPostsAfter(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + ret := _m.Called(options) var r0 *model.PostList - if rf, ok := ret.Get(0).(func(string, string, int, int) *model.PostList); ok { - r0 = rf(channelId, postId, numPosts, offset) + if rf, ok := ret.Get(0).(func(model.GetPostsOptions) *model.PostList); ok { + r0 = rf(options) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostList) @@ -421,8 +421,8 @@ func (_m *PostStore) GetPostsAfter(channelId string, postId string, numPosts int } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string, string, int, int) *model.AppError); ok { - r1 = rf(channelId, postId, numPosts, offset) + if rf, ok := ret.Get(1).(func(model.GetPostsOptions) *model.AppError); ok { + r1 = rf(options) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) @@ -457,13 +457,13 @@ func (_m *PostStore) GetPostsBatchForIndexing(startTime int64, endTime int64, li return r0, r1 } -// GetPostsBefore provides a mock function with given fields: channelId, postId, numPosts, offset -func (_m *PostStore) GetPostsBefore(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) { - ret := _m.Called(channelId, postId, numPosts, offset) +// GetPostsBefore provides a mock function with given fields: options +func (_m *PostStore) GetPostsBefore(options model.GetPostsOptions) (*model.PostList, *model.AppError) { + ret := _m.Called(options) var r0 *model.PostList - if rf, ok := ret.Get(0).(func(string, string, int, int) *model.PostList); ok { - r0 = rf(channelId, postId, numPosts, offset) + if rf, ok := ret.Get(0).(func(model.GetPostsOptions) *model.PostList); ok { + r0 = rf(options) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostList) @@ -471,8 +471,8 @@ func (_m *PostStore) GetPostsBefore(channelId string, postId string, numPosts in } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string, string, int, int) *model.AppError); ok { - r1 = rf(channelId, postId, numPosts, offset) + if rf, ok := ret.Get(1).(func(model.GetPostsOptions) *model.AppError); ok { + r1 = rf(options) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) @@ -532,13 +532,13 @@ func (_m *PostStore) GetPostsCreatedAt(channelId string, time int64) ([]*model.P return r0, r1 } -// GetPostsSince provides a mock function with given fields: channelId, time, allowFromCache -func (_m *PostStore) GetPostsSince(channelId string, time int64, allowFromCache bool) (*model.PostList, *model.AppError) { - ret := _m.Called(channelId, time, allowFromCache) +// GetPostsSince provides a mock function with given fields: options, allowFromCache +func (_m *PostStore) GetPostsSince(options model.GetPostsSinceOptions, allowFromCache bool) (*model.PostList, *model.AppError) { + ret := _m.Called(options, allowFromCache) var r0 *model.PostList - if rf, ok := ret.Get(0).(func(string, int64, bool) *model.PostList); ok { - r0 = rf(channelId, time, allowFromCache) + if rf, ok := ret.Get(0).(func(model.GetPostsSinceOptions, bool) *model.PostList); ok { + r0 = rf(options, allowFromCache) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.PostList) @@ -546,8 +546,8 @@ func (_m *PostStore) GetPostsSince(channelId string, time int64, allowFromCache } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string, int64, bool) *model.AppError); ok { - r1 = rf(channelId, time, allowFromCache) + if rf, ok := ret.Get(1).(func(model.GetPostsSinceOptions, bool) *model.AppError); ok { + r1 = rf(options, allowFromCache) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) diff --git a/store/storetest/post_store.go b/store/storetest/post_store.go index 0871607ece..13fe753f05 100644 --- a/store/storetest/post_store.go +++ b/store/storetest/post_store.go @@ -127,14 +127,14 @@ func testPostStoreGet(t *testing.T, ss store.Store) { etag2 := ss.Post().GetEtag(o1.ChannelId, false) require.Equal(t, 0, strings.Index(etag2, fmt.Sprintf("%v.%v", model.CurrentVersion, o1.UpdateAt)), "Invalid Etag") - r1, err := ss.Post().Get(o1.Id) + r1, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) require.Equal(t, r1.Posts[o1.Id].CreateAt, o1.CreateAt, "invalid returned post") - _, err = ss.Post().Get("123") + _, err = ss.Post().Get("123", false) require.NotNil(t, err, "Missing id should have failed") - _, err = ss.Post().Get("") + _, err = ss.Post().Get("", false) require.NotNil(t, err, "should fail for blank post ids") } @@ -179,15 +179,15 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { o3, err = ss.Post().Save(o3) require.Nil(t, err) - r1, err := ss.Post().Get(o1.Id) + r1, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) ro1 := r1.Posts[o1.Id] - r2, err := ss.Post().Get(o1.Id) + r2, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) ro2 := r2.Posts[o2.Id] - r3, err := ss.Post().Get(o3.Id) + r3, err := ss.Post().Get(o3.Id, false) require.Nil(t, err) ro3 := r3.Posts[o3.Id] @@ -199,7 +199,7 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { _, err = ss.Post().Update(o1a, ro1) require.Nil(t, err) - r1, err = ss.Post().Get(o1.Id) + r1, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) ro1a := r1.Posts[o1.Id] @@ -211,7 +211,7 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { _, err = ss.Post().Update(o2a, ro2) require.Nil(t, err) - r2, err = ss.Post().Get(o1.Id) + r2, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) ro2a := r2.Posts[o2.Id] @@ -223,7 +223,7 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { _, err = ss.Post().Update(o3a, ro3) require.Nil(t, err) - r3, err = ss.Post().Get(o3.Id) + r3, err = ss.Post().Get(o3.Id, false) require.Nil(t, err) ro3a := r3.Posts[o3.Id] @@ -239,7 +239,7 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { }) require.Nil(t, err) - r4, err := ss.Post().Get(o4.Id) + r4, err := ss.Post().Get(o4.Id, false) require.Nil(t, err) ro4 := r4.Posts[o4.Id] @@ -250,7 +250,7 @@ func testPostStoreUpdate(t *testing.T, ss store.Store) { _, err = ss.Post().Update(o4a, ro4) require.Nil(t, err) - r4, err = ss.Post().Get(o4.Id) + r4, err = ss.Post().Get(o4.Id, false) require.Nil(t, err) ro4a := r4.Posts[o4.Id] @@ -271,7 +271,7 @@ func testPostStoreDelete(t *testing.T, ss store.Store) { o1, err := ss.Post().Save(o1) require.Nil(t, err) - r1, err := ss.Post().Get(o1.Id) + r1, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) require.Equal(t, r1.Posts[o1.Id].CreateAt, o1.CreateAt, "invalid returned post") @@ -284,7 +284,7 @@ func testPostStoreDelete(t *testing.T, ss store.Store) { assert.Equal(t, deleteByID, actual, "Expected (*Post).Props[model.POST_PROPS_DELETE_BY] to be %v but got %v.", deleteByID, actual) - r3, err := ss.Post().Get(o1.Id) + r3, err := ss.Post().Get(o1.Id, false) require.NotNil(t, err, "Missing id should have failed - PostList %v", r3) etag2 := ss.Post().GetEtag(o1.ChannelId, false) @@ -311,10 +311,10 @@ func testPostStoreDelete1Level(t *testing.T, ss store.Store) { err = ss.Post().Delete(o1.Id, model.GetMillis(), "") require.Nil(t, err) - _, err = ss.Post().Get(o1.Id) + _, err = ss.Post().Get(o1.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o2.Id) + _, err = ss.Post().Get(o2.Id, false) require.NotNil(t, err, "Deleted id should have failed") } @@ -354,16 +354,16 @@ func testPostStoreDelete2Level(t *testing.T, ss store.Store) { err = ss.Post().Delete(o1.Id, model.GetMillis(), "") require.Nil(t, err) - _, err = ss.Post().Get(o1.Id) + _, err = ss.Post().Get(o1.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o2.Id) + _, err = ss.Post().Get(o2.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o3.Id) + _, err = ss.Post().Get(o3.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o4.Id) + _, err = ss.Post().Get(o4.Id, false) require.Nil(t, err) } @@ -394,16 +394,16 @@ func testPostStorePermDelete1Level(t *testing.T, ss store.Store) { err2 := ss.Post().PermanentDeleteByUser(o2.UserId) require.Nil(t, err2) - _, err = ss.Post().Get(o1.Id) + _, err = ss.Post().Get(o1.Id, false) require.Nil(t, err, "Deleted id shouldn't have failed") - _, err = ss.Post().Get(o2.Id) + _, err = ss.Post().Get(o2.Id, false) require.NotNil(t, err, "Deleted id should have failed") err = ss.Post().PermanentDeleteByChannel(o3.ChannelId) require.Nil(t, err) - _, err = ss.Post().Get(o3.Id) + _, err = ss.Post().Get(o3.Id, false) require.NotNil(t, err, "Deleted id should have failed") } @@ -434,13 +434,13 @@ func testPostStorePermDelete1Level2(t *testing.T, ss store.Store) { err2 := ss.Post().PermanentDeleteByUser(o1.UserId) require.Nil(t, err2) - _, err = ss.Post().Get(o1.Id) + _, err = ss.Post().Get(o1.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o2.Id) + _, err = ss.Post().Get(o2.Id, false) require.NotNil(t, err, "Deleted id should have failed") - _, err = ss.Post().Get(o3.Id) + _, err = ss.Post().Get(o3.Id, false) require.Nil(t, err, "Deleted id should have failed") } @@ -470,7 +470,7 @@ func testPostStoreGetWithChildren(t *testing.T, ss store.Store) { o3, err = ss.Post().Save(o3) require.Nil(t, err) - pl, err := ss.Post().Get(o1.Id) + pl, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) require.Len(t, pl.Posts, 3, "invalid returned post") @@ -478,7 +478,7 @@ func testPostStoreGetWithChildren(t *testing.T, ss store.Store) { dErr := ss.Post().Delete(o3.Id, model.GetMillis(), "") require.Nil(t, dErr) - pl, err = ss.Post().Get(o1.Id) + pl, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) require.Len(t, pl.Posts, 2, "invalid returned post") @@ -486,7 +486,7 @@ func testPostStoreGetWithChildren(t *testing.T, ss store.Store) { dErr = ss.Post().Delete(o2.Id, model.GetMillis(), "") require.Nil(t, dErr) - pl, err = ss.Post().Get(o1.Id) + pl, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) require.Len(t, pl.Posts, 1, "invalid returned post") @@ -548,7 +548,7 @@ func testPostStoreGetPostsWithDetails(t *testing.T, ss store.Store) { o5, err = ss.Post().Save(o5) require.Nil(t, err) - r1, err := ss.Post().GetPosts(o1.ChannelId, 0, 4, false) + r1, err := ss.Post().GetPosts(model.GetPostsOptions{ChannelId: o1.ChannelId, Page: 0, PerPage: 4}, false) require.Nil(t, err) require.Equal(t, r1.Order[0], o5.Id, "invalid order") @@ -561,7 +561,7 @@ func testPostStoreGetPostsWithDetails(t *testing.T, ss store.Store) { require.Equal(t, r1.Posts[o1.Id].Message, o1.Message, "Missing parent") - r2, err := ss.Post().GetPosts(o1.ChannelId, 0, 4, true) + r2, err := ss.Post().GetPosts(model.GetPostsOptions{ChannelId: o1.ChannelId, Page: 0, PerPage: 4}, false) require.Nil(t, err) require.Equal(t, r2.Order[0], o5.Id, "invalid order") @@ -575,7 +575,7 @@ func testPostStoreGetPostsWithDetails(t *testing.T, ss store.Store) { require.Equal(t, r2.Posts[o1.Id].Message, o1.Message, "Missing parent") // Run once to fill cache - _, err = ss.Post().GetPosts(o1.ChannelId, 0, 30, false) + _, err = ss.Post().GetPosts(model.GetPostsOptions{ChannelId: o1.ChannelId, Page: 0, PerPage: 30}, false) require.Nil(t, err) o6 := &model.Post{} @@ -585,7 +585,7 @@ func testPostStoreGetPostsWithDetails(t *testing.T, ss store.Store) { _, err = ss.Post().Save(o6) require.Nil(t, err) - r3, err := ss.Post().GetPosts(o1.ChannelId, 0, 30, false) + r3, err := ss.Post().GetPosts(model.GetPostsOptions{ChannelId: o1.ChannelId, Page: 0, PerPage: 30}, false) require.Nil(t, err) assert.Equal(t, 7, len(r3.Order)) } @@ -610,7 +610,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { } t.Run("should not return anything before the first post", func(t *testing.T) { - postList, err := ss.Post().GetPostsBefore(channelId, posts[0].Id, 10, 0) + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: posts[0].Id, Page: 0, PerPage: 10}) assert.Nil(t, err) assert.Equal(t, []string{}, postList.Order) @@ -618,7 +618,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should return posts before a post", func(t *testing.T) { - postList, err := ss.Post().GetPostsBefore(channelId, posts[5].Id, 10, 0) + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: posts[5].Id, Page: 0, PerPage: 10}) assert.Nil(t, err) assert.Equal(t, []string{posts[4].Id, posts[3].Id, posts[2].Id, posts[1].Id, posts[0].Id}, postList.Order) @@ -632,7 +632,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should limit posts before", func(t *testing.T) { - postList, err := ss.Post().GetPostsBefore(channelId, posts[5].Id, 2, 0) + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: posts[5].Id, PerPage: 2}) assert.Nil(t, err) assert.Equal(t, []string{posts[4].Id, posts[3].Id}, postList.Order) @@ -643,7 +643,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should not return anything after the last post", func(t *testing.T) { - postList, err := ss.Post().GetPostsAfter(channelId, posts[len(posts)-1].Id, 10, 0) + postList, err := ss.Post().GetPostsAfter(model.GetPostsOptions{ChannelId: channelId, PostId: posts[len(posts)-1].Id, PerPage: 10}) assert.Nil(t, err) assert.Equal(t, []string{}, postList.Order) @@ -651,7 +651,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should return posts after a post", func(t *testing.T) { - postList, err := ss.Post().GetPostsAfter(channelId, posts[5].Id, 10, 0) + postList, err := ss.Post().GetPostsAfter(model.GetPostsOptions{ChannelId: channelId, PostId: posts[5].Id, PerPage: 10}) assert.Nil(t, err) assert.Equal(t, []string{posts[9].Id, posts[8].Id, posts[7].Id, posts[6].Id}, postList.Order) @@ -664,7 +664,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should limit posts after", func(t *testing.T) { - postList, err := ss.Post().GetPostsAfter(channelId, posts[5].Id, 2, 0) + postList, err := ss.Post().GetPostsAfter(model.GetPostsOptions{ChannelId: channelId, PostId: posts[5].Id, PerPage: 2}) assert.Nil(t, err) assert.Equal(t, []string{posts[7].Id, posts[6].Id}, postList.Order) @@ -674,7 +674,6 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }, postList.Posts) }) }) - t.Run("with threads", func(t *testing.T) { channelId := model.NewId() userId := model.NewId() @@ -745,7 +744,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { post2.UpdateAt = post6.UpdateAt t.Run("should return each post and thread before a post", func(t *testing.T) { - postList, err := ss.Post().GetPostsBefore(channelId, post4.Id, 2, 0) + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: post4.Id, PerPage: 2}) assert.Nil(t, err) assert.Equal(t, []string{post3.Id, post2.Id}, postList.Order) @@ -759,7 +758,7 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }) t.Run("should return each post and the root of each thread after a post", func(t *testing.T) { - postList, err := ss.Post().GetPostsAfter(channelId, post4.Id, 2, 0) + postList, err := ss.Post().GetPostsAfter(model.GetPostsOptions{ChannelId: channelId, PostId: post4.Id, PerPage: 2}) assert.Nil(t, err) assert.Equal(t, []string{post6.Id, post5.Id}, postList.Order) @@ -771,6 +770,112 @@ func testPostStoreGetPostsBeforeAfter(t *testing.T, ss store.Store) { }, postList.Posts) }) }) + t.Run("with threads (skipFetchThreads)", func(t *testing.T) { + channelId := model.NewId() + userId := model.NewId() + + // This creates a series of posts that looks like: + // post1 + // post2 + // post3 (in response to post1) + // post4 (in response to post2) + // post5 + // post6 (in response to post2) + + post1, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + Message: "post1", + }) + require.Nil(t, err) + post1.ReplyCount = 1 + time.Sleep(time.Millisecond) + + post2, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + Message: "post2", + }) + require.Nil(t, err) + post2.ReplyCount = 2 + time.Sleep(time.Millisecond) + + post3, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + ParentId: post1.Id, + RootId: post1.Id, + Message: "post3", + }) + require.Nil(t, err) + time.Sleep(time.Millisecond) + + post4, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + RootId: post2.Id, + ParentId: post2.Id, + Message: "post4", + }) + require.Nil(t, err) + time.Sleep(time.Millisecond) + + post5, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + Message: "post5", + }) + require.Nil(t, err) + time.Sleep(time.Millisecond) + + post6, err := ss.Post().Save(&model.Post{ + ChannelId: channelId, + UserId: userId, + ParentId: post2.Id, + RootId: post2.Id, + Message: "post6", + }) + require.Nil(t, err) + + // Adding a post to a thread changes the UpdateAt timestamp of the parent post + post1.UpdateAt = post3.UpdateAt + post2.UpdateAt = post6.UpdateAt + + t.Run("should return each post and thread before a post", func(t *testing.T) { + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: post4.Id, PerPage: 2, SkipFetchThreads: true}) + assert.Nil(t, err) + + assert.Equal(t, []string{post3.Id, post2.Id}, postList.Order) + assert.Equal(t, map[string]*model.Post{ + post1.Id: post1, + post2.Id: post2, + post3.Id: post3, + }, postList.Posts) + }) + + t.Run("should return each post and thread before a post with limit", func(t *testing.T) { + postList, err := ss.Post().GetPostsBefore(model.GetPostsOptions{ChannelId: channelId, PostId: post4.Id, PerPage: 1, SkipFetchThreads: true}) + assert.Nil(t, err) + + assert.Equal(t, []string{post3.Id}, postList.Order) + assert.Equal(t, map[string]*model.Post{ + post1.Id: post1, + post3.Id: post3, + }, postList.Posts) + }) + + t.Run("should return each post and the root of each thread after a post", func(t *testing.T) { + postList, err := ss.Post().GetPostsAfter(model.GetPostsOptions{ChannelId: channelId, PostId: post4.Id, PerPage: 2, SkipFetchThreads: true}) + assert.Nil(t, err) + + assert.Equal(t, []string{post6.Id, post5.Id}, postList.Order) + assert.Equal(t, map[string]*model.Post{ + post2.Id: post2, + post5.Id: post5, + post6.Id: post6, + }, postList.Posts) + }) + }) } func testPostStoreGetPostsSince(t *testing.T, ss store.Store) { @@ -828,7 +933,7 @@ func testPostStoreGetPostsSince(t *testing.T, ss store.Store) { require.Nil(t, err) time.Sleep(time.Millisecond) - postList, err := ss.Post().GetPostsSince(channelId, post3.CreateAt, false) + postList, err := ss.Post().GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: post3.CreateAt}, false) assert.Nil(t, err) assert.Equal(t, []string{ @@ -859,7 +964,7 @@ func testPostStoreGetPostsSince(t *testing.T, ss store.Store) { require.Nil(t, err) time.Sleep(time.Millisecond) - postList, err := ss.Post().GetPostsSince(channelId, post1.CreateAt, false) + postList, err := ss.Post().GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: post1.CreateAt}, false) assert.Nil(t, err) assert.Equal(t, []string{}, postList.Order) @@ -881,12 +986,12 @@ func testPostStoreGetPostsSince(t *testing.T, ss store.Store) { time.Sleep(time.Millisecond) // Make a request that returns no results - postList, err := ss.Post().GetPostsSince(channelId, post1.CreateAt, true) + postList, err := ss.Post().GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: post1.CreateAt}, true) require.Nil(t, err) require.Equal(t, model.NewPostList(), postList) // And then ensure that it doesn't cause future requests to also return no results - postList, err = ss.Post().GetPostsSince(channelId, post1.CreateAt-1, true) + postList, err = ss.Post().GetPostsSince(model.GetPostsSinceOptions{ChannelId: channelId, Time: post1.CreateAt - 1}, true) assert.Nil(t, err) assert.Equal(t, []string{post1.Id}, postList.Order) @@ -1924,15 +2029,15 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { o3, err = ss.Post().Save(o3) require.Nil(t, err) - r1, err := ss.Post().Get(o1.Id) + r1, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) ro1 := r1.Posts[o1.Id] - r2, err := ss.Post().Get(o1.Id) + r2, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) ro2 := r2.Posts[o2.Id] - r3, err := ss.Post().Get(o3.Id) + r3, err := ss.Post().Get(o3.Id, false) require.Nil(t, err) ro3 := r3.Posts[o3.Id] @@ -1944,7 +2049,7 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { _, err = ss.Post().Overwrite(o1a) require.Nil(t, err) - r1, err = ss.Post().Get(o1.Id) + r1, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) ro1a := r1.Posts[o1.Id] @@ -1956,7 +2061,7 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { _, err = ss.Post().Overwrite(o2a) require.Nil(t, err) - r2, err = ss.Post().Get(o1.Id) + r2, err = ss.Post().Get(o1.Id, false) require.Nil(t, err) ro2a := r2.Posts[o2.Id] @@ -1968,7 +2073,7 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { _, err = ss.Post().Overwrite(o3a) require.Nil(t, err) - r3, err = ss.Post().Get(o3.Id) + r3, err = ss.Post().Get(o3.Id, false) require.Nil(t, err) ro3a := r3.Posts[o3.Id] @@ -1982,7 +2087,7 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { }) require.Nil(t, err) - r4, err := ss.Post().Get(o4.Id) + r4, err := ss.Post().Get(o4.Id, false) require.Nil(t, err) ro4 := r4.Posts[o4.Id] @@ -1993,7 +2098,7 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { _, err = ss.Post().Overwrite(o4a) require.Nil(t, err) - r4, err = ss.Post().Get(o4.Id) + r4, err = ss.Post().Get(o4.Id, false) require.Nil(t, err) ro4a := r4.Posts[o4.Id] @@ -2023,15 +2128,15 @@ func testPostStoreGetPostsByIds(t *testing.T, ss store.Store) { o3, err = ss.Post().Save(o3) require.Nil(t, err) - r1, err := ss.Post().Get(o1.Id) + r1, err := ss.Post().Get(o1.Id, false) require.Nil(t, err) ro1 := r1.Posts[o1.Id] - r2, err := ss.Post().Get(o2.Id) + r2, err := ss.Post().Get(o2.Id, false) require.Nil(t, err) ro2 := r2.Posts[o2.Id] - r3, err := ss.Post().Get(o3.Id) + r3, err := ss.Post().Get(o3.Id, false) require.Nil(t, err) ro3 := r3.Posts[o3.Id] @@ -2138,13 +2243,13 @@ func testPostStorePermanentDeleteBatch(t *testing.T, ss store.Store) { _, err = ss.Post().PermanentDeleteBatch(2000, 1000) require.Nil(t, err) - _, err = ss.Post().Get(o1.Id) + _, err = ss.Post().Get(o1.Id, false) require.NotNil(t, err, "Should have not found post 1 after purge") - _, err = ss.Post().Get(o2.Id) + _, err = ss.Post().Get(o2.Id, false) require.NotNil(t, err, "Should have not found post 2 after purge") - _, err = ss.Post().Get(o3.Id) + _, err = ss.Post().Get(o3.Id, false) require.Nil(t, err, "Should have not found post 3 after purge") } diff --git a/store/storetest/reaction_store.go b/store/storetest/reaction_store.go index b9ffd2b6c1..72625f3d83 100644 --- a/store/storetest/reaction_store.go +++ b/store/storetest/reaction_store.go @@ -43,15 +43,13 @@ func testReactionSave(t *testing.T, ss store.Store) { assert.Equal(t, saved.EmojiName, reaction1.EmojiName, "should've saved reaction emoji_name and returned it") var secondUpdateAt int64 - postList, err := ss.Post().Get(reaction1.PostId) - if err != nil { - t.Fatal(err) - } - if !postList.Posts[post.Id].HasReactions { - t.Fatal("should've set HasReactions = true on post") - } else if postList.Posts[post.Id].UpdateAt == firstUpdateAt { - t.Fatal("should've marked post as updated when HasReactions changed") - } else { + postList, err := ss.Post().Get(reaction1.PostId, false) + require.Nil(t, err) + + assert.True(t, postList.Posts[post.Id].HasReactions, "should've set HasReactions = true on post") + assert.NotEqual(t, postList.Posts[post.Id].UpdateAt, firstUpdateAt, "should've marked post as updated when HasReactions changed") + + if postList.Posts[post.Id].HasReactions && postList.Posts[post.Id].UpdateAt != firstUpdateAt { secondUpdateAt = postList.Posts[post.Id].UpdateAt } @@ -67,10 +65,8 @@ func testReactionSave(t *testing.T, ss store.Store) { _, err = ss.Reaction().Save(reaction2) require.Nil(t, err) - postList, err = ss.Post().Get(reaction2.PostId) - if err != nil { - t.Fatal(err) - } + postList, err = ss.Post().Get(reaction2.PostId, false) + require.Nil(t, err) assert.NotEqual(t, postList.Posts[post.Id].UpdateAt, secondUpdateAt, "should've marked post as updated even if HasReactions doesn't change") @@ -117,10 +113,10 @@ func testReactionDelete(t *testing.T, ss store.Store) { _, err = ss.Reaction().Save(reaction) require.Nil(t, err) - result, err := ss.Post().Get(reaction.PostId) - if err != nil { - t.Fatal(err) - } + + result, err := ss.Post().Get(reaction.PostId, false) + require.Nil(t, err) + firstUpdateAt := result.Posts[post.Id].UpdateAt _, err = ss.Reaction().Delete(reaction) @@ -131,20 +127,11 @@ func testReactionDelete(t *testing.T, ss store.Store) { assert.Empty(t, reactions, "should've deleted reaction") - if reactions, rErr := ss.Reaction().GetForPost(post.Id, false); rErr != nil { - t.Fatal(rErr) - } else if len(reactions) != 0 { - t.Fatal("should've deleted reaction") - } - postList, err := ss.Post().Get(post.Id) - if err != nil { - t.Fatal(err) - } - if postList.Posts[post.Id].HasReactions { - t.Fatal("should've set HasReactions = false on post") - } else if postList.Posts[post.Id].UpdateAt == firstUpdateAt { - t.Fatal("should mark post as updated after deleting reactions") - } + postList, err := ss.Post().Get(post.Id, false) + require.Nil(t, err) + + assert.False(t, postList.Posts[post.Id].HasReactions, "should've set HasReactions = false on post") + assert.NotEqual(t, postList.Posts[post.Id].UpdateAt, firstUpdateAt, "should mark post as updated after deleting reactions") } func testReactionGetForPost(t *testing.T, ss store.Store) { @@ -301,26 +288,17 @@ func testReactionDeleteAllWithEmojiName(t *testing.T, ss store.Store) { assert.Empty(t, returned, "should've only removed reactions with emoji name") // check that the posts are updated - postList, err := ss.Post().Get(post.Id) - if err != nil { - t.Fatal(err) - } - if !postList.Posts[post.Id].HasReactions { - t.Fatal("post should still have reactions") - } + postList, err := ss.Post().Get(post.Id, false) + require.Nil(t, err) + assert.True(t, postList.Posts[post.Id].HasReactions, "post should still have reactions") - postList, err = ss.Post().Get(post2.Id) - if err != nil { - t.Fatal(err) - } - if !postList.Posts[post2.Id].HasReactions { - t.Fatal("post should still have reactions") - } + postList, err = ss.Post().Get(post2.Id, false) + require.Nil(t, err) + assert.True(t, postList.Posts[post2.Id].HasReactions, "post should still have reactions") - postList, err = ss.Post().Get(post3.Id) - if err != nil { - t.Fatal(err) - } + postList, err = ss.Post().Get(post3.Id, false) + require.Nil(t, err) + assert.False(t, postList.Posts[post3.Id].HasReactions, "post shouldn't have reactions any more") } diff --git a/store/timer_layer.go b/store/timer_layer.go index ad1f4cffe9..f465fe435c 100644 --- a/store/timer_layer.go +++ b/store/timer_layer.go @@ -3645,6 +3645,22 @@ func (s *TimerLayerOAuthStore) UpdateApp(app *model.OAuthApp) (*model.OAuthApp, return resultVar0, resultVar1 } +func (s *TimerLayerPluginStore) CompareAndDelete(keyVal *model.PluginKeyValue, oldValue []byte) (bool, *model.AppError) { + start := timemodule.Now() + + resultVar0, resultVar1 := s.PluginStore.CompareAndDelete(keyVal, oldValue) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if resultVar1 == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PluginStore.CompareAndDelete", success, elapsed) + } + return resultVar0, resultVar1 +} + func (s *TimerLayerPluginStore) CompareAndSet(keyVal *model.PluginKeyValue, oldValue []byte) (bool, *model.AppError) { start := timemodule.Now() @@ -3852,10 +3868,10 @@ func (s *TimerLayerPostStore) Delete(postId string, time int64, deleteByID strin return resultVar0 } -func (s *TimerLayerPostStore) Get(id string) (*model.PostList, *model.AppError) { +func (s *TimerLayerPostStore) Get(id string, skipFetchThreads bool) (*model.PostList, *model.AppError) { start := timemodule.Now() - resultVar0, resultVar1 := s.PostStore.Get(id) + resultVar0, resultVar1 := s.PostStore.Get(id, skipFetchThreads) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -4044,10 +4060,10 @@ func (s *TimerLayerPostStore) GetPostIdBeforeTime(channelId string, time int64) return resultVar0, resultVar1 } -func (s *TimerLayerPostStore) GetPosts(channelId string, offset int, limit int, allowFromCache bool) (*model.PostList, *model.AppError) { +func (s *TimerLayerPostStore) GetPosts(options model.GetPostsOptions, allowFromCache bool) (*model.PostList, *model.AppError) { start := timemodule.Now() - resultVar0, resultVar1 := s.PostStore.GetPosts(channelId, offset, limit, allowFromCache) + resultVar0, resultVar1 := s.PostStore.GetPosts(options, allowFromCache) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -4060,10 +4076,10 @@ func (s *TimerLayerPostStore) GetPosts(channelId string, offset int, limit int, return resultVar0, resultVar1 } -func (s *TimerLayerPostStore) GetPostsAfter(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) { +func (s *TimerLayerPostStore) GetPostsAfter(options model.GetPostsOptions) (*model.PostList, *model.AppError) { start := timemodule.Now() - resultVar0, resultVar1 := s.PostStore.GetPostsAfter(channelId, postId, numPosts, offset) + resultVar0, resultVar1 := s.PostStore.GetPostsAfter(options) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -4092,10 +4108,10 @@ func (s *TimerLayerPostStore) GetPostsBatchForIndexing(startTime int64, endTime return resultVar0, resultVar1 } -func (s *TimerLayerPostStore) GetPostsBefore(channelId string, postId string, numPosts int, offset int) (*model.PostList, *model.AppError) { +func (s *TimerLayerPostStore) GetPostsBefore(options model.GetPostsOptions) (*model.PostList, *model.AppError) { start := timemodule.Now() - resultVar0, resultVar1 := s.PostStore.GetPostsBefore(channelId, postId, numPosts, offset) + resultVar0, resultVar1 := s.PostStore.GetPostsBefore(options) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -4140,10 +4156,10 @@ func (s *TimerLayerPostStore) GetPostsCreatedAt(channelId string, time int64) ([ return resultVar0, resultVar1 } -func (s *TimerLayerPostStore) GetPostsSince(channelId string, time int64, allowFromCache bool) (*model.PostList, *model.AppError) { +func (s *TimerLayerPostStore) GetPostsSince(options model.GetPostsSinceOptions, allowFromCache bool) (*model.PostList, *model.AppError) { start := timemodule.Now() - resultVar0, resultVar1 := s.PostStore.GetPostsSince(channelId, time, allowFromCache) + resultVar0, resultVar1 := s.PostStore.GetPostsSince(options, allowFromCache) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil {