diff --git a/server/channels/app/group.go b/server/channels/app/group.go index 35d7ae9b7a..9c01468ee7 100644 --- a/server/channels/app/group.go +++ b/server/channels/app/group.go @@ -9,6 +9,7 @@ import ( "net/http" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/utils" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -25,7 +26,10 @@ func (a *App) GetGroup(id string, opts *model.GetGroupOpts, viewRestrictions *mo } if opts != nil && opts.IncludeMemberIDs { - users, err := a.Srv().Store().Group().GetMemberUsers(id) + perPage := 100 + users, err := utils.Pager(func(page int) ([]*model.User, error) { + return a.Srv().Store().Group().GetMemberUsersPage(id, page, perPage, viewRestrictions) + }, perPage) if err != nil { return nil, model.NewAppError("GetGroup", "app.member_count", nil, "", http.StatusInternalServerError).Wrap(err) } diff --git a/server/channels/app/group_test.go b/server/channels/app/group_test.go index 121e76ccbd..264e750dee 100644 --- a/server/channels/app/group_test.go +++ b/server/channels/app/group_test.go @@ -12,27 +12,82 @@ import ( "github.com/mattermost/mattermost/server/public/model" ) +// TestGetGroup tests basic group retrieval and verifies that ViewUsersRestrictions +// are properly applied when fetching member IDs to prevent guest users from +// enumerating user IDs they shouldn't have access to. func TestGetGroup(t *testing.T) { mainHelper.Parallel(t) - th := Setup(t) + th := Setup(t).InitBasic() defer th.TearDown() - group := th.CreateGroup() - group, err := th.App.GetGroup(group.Id, nil, nil) - require.Nil(t, err) - require.NotNil(t, group) + t.Run("basic retrieval", func(t *testing.T) { + group := th.CreateGroup() - nilGroup, err := th.App.GetGroup(model.NewId(), nil, nil) - require.NotNil(t, err) - require.Nil(t, nilGroup) + g, err := th.App.GetGroup(group.Id, nil, nil) + require.Nil(t, err) + require.NotNil(t, g) - group, err = th.App.GetGroup(group.Id, &model.GetGroupOpts{IncludeMemberCount: false}, nil) - require.Nil(t, err) - require.Nil(t, group.MemberCount) + nilGroup, err := th.App.GetGroup(model.NewId(), nil, nil) + require.NotNil(t, err) + require.Nil(t, nilGroup) + }) - group, err = th.App.GetGroup(group.Id, &model.GetGroupOpts{IncludeMemberCount: true}, nil) - require.Nil(t, err) - require.NotNil(t, group.MemberCount) + t.Run("include member count", func(t *testing.T) { + group := th.CreateGroup() + + g, err := th.App.GetGroup(group.Id, &model.GetGroupOpts{IncludeMemberCount: false}, nil) + require.Nil(t, err) + require.Nil(t, g.MemberCount) + + g, err = th.App.GetGroup(group.Id, &model.GetGroupOpts{IncludeMemberCount: true}, nil) + require.Nil(t, err) + require.NotNil(t, g.MemberCount) + }) + + t.Run("member IDs respect view restrictions", func(t *testing.T) { + user1 := th.CreateUser() + user2 := th.CreateUser() + user3 := th.CreateUser() + + team := th.CreateTeam() + channel := th.CreateChannel(th.Context, team) + + th.LinkUserToTeam(user1, team) + th.AddUserToChannel(user1, channel) + + id := model.NewId() + groupWithUserIds := &model.GroupWithUserIds{ + Group: model.Group{ + DisplayName: "dn_" + id, + Name: model.NewPointer("name" + id), + Source: model.GroupSourceCustom, + AllowReference: true, + }, + UserIds: []string{user1.Id, user2.Id, user3.Id}, + } + group, err := th.App.CreateGroupWithUserIds(groupWithUserIds) + require.Nil(t, err) + + opts := &model.GetGroupOpts{IncludeMemberIDs: true} + + g, appErr := th.App.GetGroup(group.Id, opts, nil) + require.Nil(t, appErr) + assert.Len(t, g.MemberIDs, 3) + + g, appErr = th.App.GetGroup(group.Id, opts, &model.ViewUsersRestrictions{Channels: []string{channel.Id}}) + require.Nil(t, appErr) + assert.Len(t, g.MemberIDs, 1) + assert.Contains(t, g.MemberIDs, user1.Id) + + g, appErr = th.App.GetGroup(group.Id, opts, &model.ViewUsersRestrictions{Teams: []string{team.Id}}) + require.Nil(t, appErr) + assert.Len(t, g.MemberIDs, 1) + assert.Contains(t, g.MemberIDs, user1.Id) + + g, appErr = th.App.GetGroup(group.Id, opts, &model.ViewUsersRestrictions{Channels: []string{}, Teams: []string{}}) + require.Nil(t, appErr) + assert.Empty(t, g.MemberIDs) + }) } func TestGetGroupByRemoteID(t *testing.T) { diff --git a/server/public/utils/page.go b/server/public/utils/page.go new file mode 100644 index 0000000000..c9e000036e --- /dev/null +++ b/server/public/utils/page.go @@ -0,0 +1,43 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package utils + +// Pager fetches all items from a paginated API. +// Pager is a generic function that fetches and aggregates paginated data. +// It takes a fetch function and a perPage parameter as arguments. +// +// The fetch function is responsible for retrieving a slice of items of type T +// for a given page number. It returns the fetched items and an error, if any. +// Ideally a developer may want to use a closure to create a fetch function. +// +// The perPage parameter specifies the number of items to fetch per page. +// +// Example usage: +// +// items, err := Pager(fetchFunc, 10) +// if err != nil { +// // handle error +// } +// // process items +func Pager[T any](fetch func(page int) ([]T, error), perPage int) ([]T, error) { + var list []T + var page int + + for { + fetched, err := fetch(page) + if err != nil { + return list, err + } + + list = append(list, fetched...) + + if len(fetched) < perPage { + break + } + + page++ + } + + return list, nil +} diff --git a/server/public/utils/page_test.go b/server/public/utils/page_test.go new file mode 100644 index 0000000000..e5a8adbe48 --- /dev/null +++ b/server/public/utils/page_test.go @@ -0,0 +1,69 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package utils + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestPager(t *testing.T) { + tests := []struct { + name string + fetch func(page int) ([]int, error) + perPage int + expected []int + expectErr bool + }{ + { + name: "successful fetch", + fetch: func(page int) ([]int, error) { + if page > 2 { + return nil, nil + } + return []int{page*10 + 1, page*10 + 2, page*10 + 3}, nil + }, + perPage: 3, + expected: []int{1, 2, 3, 11, 12, 13, 21, 22, 23}, + }, + { + name: "fetch with error", + fetch: func(page int) ([]int, error) { + if page == 1 { + return nil, errors.New("fetch error") + } + return []int{page*10 + 1, page*10 + 2, page*10 + 3}, nil + }, + perPage: 3, + expected: []int{1, 2, 3}, + expectErr: true, + }, + { + name: "fetch with fewer items than perPage", + fetch: func(page int) ([]int, error) { + if page > 0 { + return []int{11, 12}, nil + } + return []int{page*10 + 1, page*10 + 2, page*10 + 3}, nil + }, + perPage: 3, + expected: []int{1, 2, 3, 11, 12}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result, err := Pager(tt.fetch, tt.perPage) + if tt.expectErr { + assert.Error(t, err) + assert.Equal(t, tt.expected, result) + } else { + assert.NoError(t, err) + assert.Equal(t, tt.expected, result) + } + }) + } +}