diff --git a/api4/post.go b/api4/post.go index 5a55771e8d..3ee2fa8c13 100644 --- a/api4/post.go +++ b/api4/post.go @@ -529,11 +529,9 @@ func getPostThread(c *Context, w http.ResponseWriter, r *http.Request) { } fromPost := r.URL.Query().Get("fromPost") - // Either both have to be set, or none have to be set. - // Setting one and not setting the other is an error. - if (fromPost == "" && fromCreateAt != 0) || (fromPost != "" && fromCreateAt == 0) { - c.SetInvalidParam("fromPost/fromCreateAt") - return + // Either only fromCreateAt must be set, or both fromPost and fromCreateAt must be set + if fromPost != "" && fromCreateAt == 0 { + c.SetInvalidParam("if fromPost is set, then fromCreatAt must also be set") } direction := "" diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index 2368dded54..a9761981da 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -604,23 +604,34 @@ func (s *SqlPostStore) getPostWithCollapsedThreads(id, userID string, opts model query = query.OrderBy("CreateAt " + sort + ", Id " + sort) } - if opts.FromPost != "" && opts.FromCreateAt != 0 { + if opts.FromCreateAt != 0 { if opts.Direction == "down" { - query = query.Where(sq.Or{ - sq.Gt{"Posts.CreateAt": opts.FromCreateAt}, - sq.And{ - sq.Eq{"Posts.CreateAt": opts.FromCreateAt}, - sq.Gt{"Posts.Id": opts.FromPost}, - }, - }) + direction := sq.Gt{"Posts.CreateAt": opts.FromCreateAt} + if opts.FromPost != "" { + query = query.Where(sq.Or{ + direction, + sq.And{ + sq.Eq{"Posts.CreateAt": opts.FromCreateAt}, + sq.Gt{"Posts.Id": opts.FromPost}, + }, + }) + } else { + query = query.Where(direction) + } } else { - query = query.Where(sq.Or{ - sq.Lt{"Posts.CreateAt": opts.FromCreateAt}, - sq.And{ - sq.Eq{"Posts.CreateAt": opts.FromCreateAt}, - sq.Lt{"Posts.Id": opts.FromPost}, - }, - }) + direction := sq.Lt{"Posts.CreateAt": opts.FromCreateAt} + if opts.FromPost != "" { + query = query.Where(sq.Or{ + direction, + sq.And{ + sq.Eq{"Posts.CreateAt": opts.FromCreateAt}, + sq.Lt{"Posts.Id": opts.FromPost}, + }, + }) + + } else { + query = query.Where(direction) + } } } @@ -715,23 +726,34 @@ func (s *SqlPostStore) Get(ctx context.Context, id string, opts model.GetPostsOp query = query.OrderBy("CreateAt " + sort + ", Id " + sort) } - if opts.FromPost != "" && opts.FromCreateAt != 0 { + if opts.FromCreateAt != 0 { if opts.Direction == "down" { - query = query.Where(sq.Or{ - sq.Gt{"p.CreateAt": opts.FromCreateAt}, - sq.And{ - sq.Eq{"p.CreateAt": opts.FromCreateAt}, - sq.Gt{"p.Id": opts.FromPost}, - }, - }) + direction := sq.Gt{"p.CreateAt": opts.FromCreateAt} + if opts.FromPost != "" { + query = query.Where(sq.Or{ + direction, + sq.And{ + sq.Eq{"p.CreateAt": opts.FromCreateAt}, + sq.Gt{"p.Id": opts.FromPost}, + }, + }) + } else { + query = query.Where(direction) + } } else { - query = query.Where(sq.Or{ - sq.Lt{"p.CreateAt": opts.FromCreateAt}, - sq.And{ - sq.Eq{"p.CreateAt": opts.FromCreateAt}, - sq.Lt{"p.Id": opts.FromPost}, - }, - }) + direction := sq.Lt{"p.CreateAt": opts.FromCreateAt} + if opts.FromPost != "" { + query = query.Where(sq.Or{ + direction, + sq.And{ + sq.Eq{"p.CreateAt": opts.FromCreateAt}, + sq.Lt{"p.Id": opts.FromPost}, + }, + }) + + } else { + query = query.Where(direction) + } } } diff --git a/store/storetest/post_store.go b/store/storetest/post_store.go index 1d98dc9a31..5f095f576a 100644 --- a/store/storetest/post_store.go +++ b/store/storetest/post_store.go @@ -567,7 +567,7 @@ func testPostStoreGetForThread(t *testing.T, ss store.Store) { require.NoError(t, err) _, err = ss.Post().Save(&model.Post{ChannelId: o1.ChannelId, UserId: model.NewId(), Message: NewTestId(), RootId: o1.Id}) require.NoError(t, err) - _, err = ss.Post().Save(&model.Post{ChannelId: o1.ChannelId, UserId: model.NewId(), Message: NewTestId(), RootId: o1.Id}) + m1, err := ss.Post().Save(&model.Post{ChannelId: o1.ChannelId, UserId: model.NewId(), Message: NewTestId(), RootId: o1.Id}) require.NoError(t, err) _, err = ss.Post().Save(&model.Post{ChannelId: o1.ChannelId, UserId: model.NewId(), Message: NewTestId(), RootId: o1.Id}) require.NoError(t, err) @@ -615,6 +615,20 @@ func testPostStoreGetForThread(t *testing.T, ss store.Store) { assert.LessOrEqual(t, r1.Posts[r1.Order[1]].CreateAt, firstPostCreateAt) assert.False(t, r1.HasNext) + // Only with CreateAt + opts = model.GetPostsOptions{ + CollapsedThreads: false, + PerPage: 1, + Direction: "up", + FromCreateAt: m1.CreateAt, + SkipFetchThreads: false, + } + r1, err = ss.Post().Get(context.Background(), o1.Id, opts, o1.UserId) + require.NoError(t, err) + assert.Len(t, r1.Order, 2) // including the root post + assert.LessOrEqual(t, r1.Posts[r1.Order[1]].CreateAt, m1.CreateAt) + assert.True(t, r1.HasNext) + // Non-CRT mode opts = model.GetPostsOptions{ CollapsedThreads: false, @@ -659,6 +673,20 @@ func testPostStoreGetForThread(t *testing.T, ss store.Store) { assert.Len(t, r1.Order, 3) // including the root post assert.LessOrEqual(t, r1.Posts[r1.Order[1]].CreateAt, firstPostCreateAt) assert.False(t, r1.HasNext) + + // Only with CreateAt + opts = model.GetPostsOptions{ + CollapsedThreads: false, + PerPage: 1, + Direction: "down", + FromCreateAt: m1.CreateAt, + SkipFetchThreads: false, + } + r1, err = ss.Post().Get(context.Background(), o1.Id, opts, o1.UserId) + require.NoError(t, err) + assert.Len(t, r1.Order, 2) // including the root post + assert.GreaterOrEqual(t, r1.Posts[r1.Order[1]].CreateAt, m1.CreateAt) + assert.True(t, r1.HasNext) }) }