From 0ec0616d892f7a0e97aaa937644afcae424c998d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Espino?= Date: Wed, 31 Jul 2019 18:34:03 +0200 Subject: [PATCH] Restricting team stats using the VIEW_MEMBERS restrictions (#11694) * Restricting team stats using the VIEW_MEMBERS restrictions * Adding tests * fixing tests --- api4/team.go | 8 ++- app/plugin_api.go | 2 +- app/team.go | 8 +-- app/team_test.go | 64 ++++++++++++++++++++--- store/sqlstore/team_store.go | 84 ++++++++++++++++++++++-------- store/store.go | 4 +- store/storetest/mocks/TeamStore.go | 28 +++++----- store/storetest/team_store.go | 16 +++--- 8 files changed, 155 insertions(+), 59 deletions(-) diff --git a/api4/team.go b/api4/team.go index 67a06aed26..73c3e8d5ba 100644 --- a/api4/team.go +++ b/api4/team.go @@ -661,7 +661,13 @@ func getTeamStats(c *Context, w http.ResponseWriter, r *http.Request) { return } - stats, err := c.App.GetTeamStats(c.Params.TeamId) + restrictions, err := c.App.GetViewUsersRestrictions(c.App.Session.UserId) + if err != nil { + c.Err = err + return + } + + stats, err := c.App.GetTeamStats(c.Params.TeamId, restrictions) if err != nil { c.Err = err return diff --git a/app/plugin_api.go b/app/plugin_api.go index 3794d0cc47..da9dfef180 100644 --- a/app/plugin_api.go +++ b/app/plugin_api.go @@ -189,7 +189,7 @@ func (api *PluginAPI) UpdateTeamMemberRoles(teamId, userId, newRoles string) (*m } func (api *PluginAPI) GetTeamStats(teamId string) (*model.TeamStats, *model.AppError) { - return api.app.GetTeamStats(teamId) + return api.app.GetTeamStats(teamId, nil) } func (api *PluginAPI) CreateUser(user *model.User) (*model.User, *model.AppError) { diff --git a/app/team.go b/app/team.go index b2d5738aa7..7bb632ee4c 100644 --- a/app/team.go +++ b/app/team.go @@ -561,7 +561,7 @@ func (a *App) joinUserToTeam(team *model.Team, user *model.User) (*model.TeamMem return rtm, true, nil } - membersCount, err := a.Srv.Store.Team().GetActiveMemberCount(tm.TeamId) + membersCount, err := a.Srv.Store.Team().GetActiveMemberCount(tm.TeamId, nil) if err != nil { return nil, false, err } @@ -1238,16 +1238,16 @@ func (a *App) RestoreTeam(teamId string) *model.AppError { return nil } -func (a *App) GetTeamStats(teamId string) (*model.TeamStats, *model.AppError) { +func (a *App) GetTeamStats(teamId string, restrictions *model.ViewUsersRestrictions) (*model.TeamStats, *model.AppError) { tchan := make(chan store.StoreResult, 1) go func() { - totalMemberCount, err := a.Srv.Store.Team().GetTotalMemberCount(teamId) + totalMemberCount, err := a.Srv.Store.Team().GetTotalMemberCount(teamId, restrictions) tchan <- store.StoreResult{Data: totalMemberCount, Err: err} close(tchan) }() achan := make(chan store.StoreResult, 1) go func() { - memberCount, err := a.Srv.Store.Team().GetActiveMemberCount(teamId) + memberCount, err := a.Srv.Store.Team().GetActiveMemberCount(teamId, restrictions) achan <- store.StoreResult{Data: memberCount, Err: err} close(achan) }() diff --git a/app/team_test.go b/app/team_test.go index 3da47b0c08..fe150f1cce 100644 --- a/app/team_test.go +++ b/app/team_test.go @@ -816,12 +816,64 @@ func TestGetTeamStats(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id) - require.Nil(t, err) - require.NotNil(t, teamStats) - members, err := th.App.GetTeamMembers(th.BasicTeam.Id, 0, 5, nil) - require.Nil(t, err) - assert.Equal(t, int64(len(members)), teamStats.TotalMemberCount) + t.Run("without view restrictions", func(t *testing.T) { + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, nil) + require.Nil(t, err) + require.NotNil(t, teamStats) + members, err := th.App.GetTeamMembers(th.BasicTeam.Id, 0, 5, nil) + require.Nil(t, err) + assert.Equal(t, int64(len(members)), teamStats.TotalMemberCount) + assert.Equal(t, int64(len(members)), teamStats.ActiveMemberCount) + }) + + t.Run("with view restrictions by this team", func(t *testing.T) { + restrictions := &model.ViewUsersRestrictions{Teams: []string{th.BasicTeam.Id}} + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, restrictions) + require.Nil(t, err) + require.NotNil(t, teamStats) + members, err := th.App.GetTeamMembers(th.BasicTeam.Id, 0, 5, nil) + require.Nil(t, err) + assert.Equal(t, int64(len(members)), teamStats.TotalMemberCount) + assert.Equal(t, int64(len(members)), teamStats.ActiveMemberCount) + }) + + t.Run("with view restrictions by valid channel", func(t *testing.T) { + restrictions := &model.ViewUsersRestrictions{Teams: []string{}, Channels: []string{th.BasicChannel.Id}} + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, restrictions) + require.Nil(t, err) + require.NotNil(t, teamStats) + members, err := th.App.GetChannelMembersPage(th.BasicChannel.Id, 0, 5) + require.Nil(t, err) + assert.Equal(t, int64(len(*members)), teamStats.TotalMemberCount) + assert.Equal(t, int64(len(*members)), teamStats.ActiveMemberCount) + }) + + t.Run("with view restrictions to not see anything", func(t *testing.T) { + restrictions := &model.ViewUsersRestrictions{Teams: []string{}, Channels: []string{}} + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, restrictions) + require.Nil(t, err) + require.NotNil(t, teamStats) + assert.Equal(t, int64(0), teamStats.TotalMemberCount) + assert.Equal(t, int64(0), teamStats.ActiveMemberCount) + }) + + t.Run("with view restrictions by other team", func(t *testing.T) { + restrictions := &model.ViewUsersRestrictions{Teams: []string{"other-team-id"}} + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, restrictions) + require.Nil(t, err) + require.NotNil(t, teamStats) + assert.Equal(t, int64(0), teamStats.TotalMemberCount) + assert.Equal(t, int64(0), teamStats.ActiveMemberCount) + }) + + t.Run("with view restrictions by not-existing channel", func(t *testing.T) { + restrictions := &model.ViewUsersRestrictions{Teams: []string{}, Channels: []string{"test"}} + teamStats, err := th.App.GetTeamStats(th.BasicTeam.Id, restrictions) + require.Nil(t, err) + require.NotNil(t, teamStats) + assert.Equal(t, int64(0), teamStats.TotalMemberCount) + assert.Equal(t, int64(0), teamStats.ActiveMemberCount) + }) } func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { diff --git a/store/sqlstore/team_store.go b/store/sqlstore/team_store.go index aac9ef78da..cd15ca0f51 100644 --- a/store/sqlstore/team_store.go +++ b/store/sqlstore/team_store.go @@ -590,35 +590,43 @@ func (s SqlTeamStore) GetMembers(teamId string, offset int, limit int, restricti return dbMembers.ToModel(), nil } -func (s SqlTeamStore) GetTotalMemberCount(teamId string) (int64, *model.AppError) { - count, err := s.GetReplica().SelectInt(` - SELECT - count(*) - FROM - TeamMembers, - Users - WHERE - TeamMembers.UserId = Users.Id - AND TeamMembers.TeamId = :TeamId - AND TeamMembers.DeleteAt = 0`, map[string]interface{}{"TeamId": teamId}) +func (s SqlTeamStore) GetTotalMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) { + query := s.getQueryBuilder(). + Select("count(DISTINCT TeamMembers.UserId)"). + From("TeamMembers, Users"). + Where("TeamMembers.DeleteAt = 0"). + Where("TeamMembers.UserId = Users.Id"). + Where(sq.Eq{"TeamMembers.TeamId": teamId}) + + query = applyTeamMemberViewRestrictionsFilterForStats(query, teamId, restrictions) + queryString, args, err := query.ToSql() + if err != nil { + return int64(0), model.NewAppError("SqlTeamStore.GetTotalMemberCount", "store.sql_team.get_member_count.app_error", nil, err.Error(), http.StatusInternalServerError) + } + + count, err := s.GetReplica().SelectInt(queryString, args...) if err != nil { return int64(0), model.NewAppError("SqlTeamStore.GetTotalMemberCount", "store.sql_team.get_member_count.app_error", nil, "teamId="+teamId+" "+err.Error(), http.StatusInternalServerError) } return count, nil } -func (s SqlTeamStore) GetActiveMemberCount(teamId string) (int64, *model.AppError) { - count, err := s.GetReplica().SelectInt(` - SELECT - count(*) - FROM - TeamMembers, - Users - WHERE - TeamMembers.UserId = Users.Id - AND TeamMembers.TeamId = :TeamId - AND TeamMembers.DeleteAt = 0 - AND Users.DeleteAt = 0`, map[string]interface{}{"TeamId": teamId}) +func (s SqlTeamStore) GetActiveMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) { + query := s.getQueryBuilder(). + Select("count(DISTINCT TeamMembers.UserId)"). + From("TeamMembers, Users"). + Where("TeamMembers.DeleteAt = 0"). + Where("TeamMembers.UserId = Users.Id"). + Where("Users.DeleteAt = 0"). + Where(sq.Eq{"TeamMembers.TeamId": teamId}) + + query = applyTeamMemberViewRestrictionsFilterForStats(query, teamId, restrictions) + queryString, args, err := query.ToSql() + if err != nil { + return 0, model.NewAppError("SqlTeamStore.GetActiveMemberCount", "store.sql_team.get_active_member_count.app_error", nil, err.Error(), http.StatusInternalServerError) + } + + count, err := s.GetReplica().SelectInt(queryString, args...) if err != nil { return 0, model.NewAppError("SqlTeamStore.GetActiveMemberCount", "store.sql_team.get_active_member_count.app_error", nil, "teamId="+teamId+" "+err.Error(), http.StatusInternalServerError) } @@ -1057,3 +1065,33 @@ func applyTeamMemberViewRestrictionsFilter(query sq.SelectBuilder, teamId string return resultQuery.Distinct() } + +func applyTeamMemberViewRestrictionsFilterForStats(query sq.SelectBuilder, teamId string, restrictions *model.ViewUsersRestrictions) sq.SelectBuilder { + if restrictions == nil { + return query + } + + // If you have no access to teams or channels, return and empty result. + if restrictions.Teams != nil && len(restrictions.Teams) == 0 && restrictions.Channels != nil && len(restrictions.Channels) == 0 { + return query.Where("1 = 0") + } + + teams := make([]interface{}, len(restrictions.Teams)) + for i, v := range restrictions.Teams { + teams[i] = v + } + channels := make([]interface{}, len(restrictions.Channels)) + for i, v := range restrictions.Channels { + channels[i] = v + } + + resultQuery := query + if restrictions.Teams != nil && len(restrictions.Teams) > 0 { + resultQuery = resultQuery.Join(fmt.Sprintf("TeamMembers rtm ON ( rtm.UserId = Users.Id AND rtm.DeleteAt = 0 AND rtm.TeamId IN (%s))", sq.Placeholders(len(teams))), teams...) + } + if restrictions.Channels != nil && len(restrictions.Channels) > 0 { + resultQuery = resultQuery.Join(fmt.Sprintf("ChannelMembers rcm ON ( rcm.UserId = Users.Id AND rcm.ChannelId IN (%s))", sq.Placeholders(len(channels))), channels...) + } + + return resultQuery +} diff --git a/store/store.go b/store/store.go index cbdb33e061..f4ec633061 100644 --- a/store/store.go +++ b/store/store.go @@ -77,8 +77,8 @@ type TeamStore interface { GetMember(teamId string, userId string) (*model.TeamMember, *model.AppError) GetMembers(teamId string, offset int, limit int, restrictions *model.ViewUsersRestrictions) ([]*model.TeamMember, *model.AppError) GetMembersByIds(teamId string, userIds []string, restrictions *model.ViewUsersRestrictions) ([]*model.TeamMember, *model.AppError) - GetTotalMemberCount(teamId string) (int64, *model.AppError) - GetActiveMemberCount(teamId string) (int64, *model.AppError) + GetTotalMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) + GetActiveMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) GetTeamsForUser(userId string) ([]*model.TeamMember, *model.AppError) GetTeamsForUserWithPagination(userId string, page, perPage int) ([]*model.TeamMember, *model.AppError) GetChannelUnreadsForAllTeams(excludeTeamId, userId string) ([]*model.ChannelUnread, *model.AppError) diff --git a/store/storetest/mocks/TeamStore.go b/store/storetest/mocks/TeamStore.go index a5c11b817c..e490238c37 100644 --- a/store/storetest/mocks/TeamStore.go +++ b/store/storetest/mocks/TeamStore.go @@ -104,20 +104,20 @@ func (_m *TeamStore) Get(id string) (*model.Team, *model.AppError) { return r0, r1 } -// GetActiveMemberCount provides a mock function with given fields: teamId -func (_m *TeamStore) GetActiveMemberCount(teamId string) (int64, *model.AppError) { - ret := _m.Called(teamId) +// GetActiveMemberCount provides a mock function with given fields: teamId, restrictions +func (_m *TeamStore) GetActiveMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) { + ret := _m.Called(teamId, restrictions) var r0 int64 - if rf, ok := ret.Get(0).(func(string) int64); ok { - r0 = rf(teamId) + if rf, ok := ret.Get(0).(func(string, *model.ViewUsersRestrictions) int64); ok { + r0 = rf(teamId, restrictions) } else { r0 = ret.Get(0).(int64) } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { - r1 = rf(teamId) + if rf, ok := ret.Get(1).(func(string, *model.ViewUsersRestrictions) *model.AppError); ok { + r1 = rf(teamId, restrictions) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) @@ -602,20 +602,20 @@ func (_m *TeamStore) GetTeamsForUserWithPagination(userId string, page int, perP return r0, r1 } -// GetTotalMemberCount provides a mock function with given fields: teamId -func (_m *TeamStore) GetTotalMemberCount(teamId string) (int64, *model.AppError) { - ret := _m.Called(teamId) +// GetTotalMemberCount provides a mock function with given fields: teamId, restrictions +func (_m *TeamStore) GetTotalMemberCount(teamId string, restrictions *model.ViewUsersRestrictions) (int64, *model.AppError) { + ret := _m.Called(teamId, restrictions) var r0 int64 - if rf, ok := ret.Get(0).(func(string) int64); ok { - r0 = rf(teamId) + if rf, ok := ret.Get(0).(func(string, *model.ViewUsersRestrictions) int64); ok { + r0 = rf(teamId, restrictions) } else { r0 = ret.Get(0).(int64) } var r1 *model.AppError - if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { - r1 = rf(teamId) + if rf, ok := ret.Get(1).(func(string, *model.ViewUsersRestrictions) *model.AppError); ok { + r1 = rf(teamId, restrictions) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) diff --git a/store/storetest/team_store.go b/store/storetest/team_store.go index 11b9b4229c..edf9fd0f55 100644 --- a/store/storetest/team_store.go +++ b/store/storetest/team_store.go @@ -903,7 +903,7 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { }(userIds[i]) } - if totalMemberCount, err := ss.Team().GetTotalMemberCount(team.Id); err != nil { + if totalMemberCount, err := ss.Team().GetTotalMemberCount(team.Id, nil); err != nil { t.Fatal(err) } else if int(totalMemberCount) != maxUsersPerTeam { t.Fatalf("should start with 5 team members, had %v instead", totalMemberCount) @@ -926,7 +926,7 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { t.Fatal("shouldn't be able to save member when at maximum members per team") } - if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id); teamErr != nil { + if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id, nil); teamErr != nil { t.Fatal(teamErr) } else if int(totalMemberCount) != maxUsersPerTeam { t.Fatalf("should still have 5 team members, had %v instead", totalMemberCount) @@ -941,7 +941,7 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { panic(teamErr) } - if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id); teamErr != nil { + if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id, nil); teamErr != nil { t.Fatal(teamErr) } else if int(totalMemberCount) != maxUsersPerTeam-1 { t.Fatalf("should now only have 4 team members, had %v instead", totalMemberCount) @@ -955,7 +955,7 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { }(newUserId) } - if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id); teamErr != nil { + if totalMemberCount, teamErr := ss.Team().GetTotalMemberCount(team.Id, nil); teamErr != nil { t.Fatal(teamErr) } else if int(totalMemberCount) != maxUsersPerTeam { t.Fatalf("should have 5 team members again, had %v instead", totalMemberCount) @@ -1117,7 +1117,7 @@ func testTeamStoreMemberCount(t *testing.T, ss store.Store) { require.Nil(t, err) var totalMemberCount int64 - if totalMemberCount, err = ss.Team().GetTotalMemberCount(teamId1); err != nil { + if totalMemberCount, err = ss.Team().GetTotalMemberCount(teamId1, nil); err != nil { t.Fatal(err) } else { if totalMemberCount != 2 { @@ -1126,7 +1126,7 @@ func testTeamStoreMemberCount(t *testing.T, ss store.Store) { } var result int64 - if result, err = ss.Team().GetActiveMemberCount(teamId1); err != nil { + if result, err = ss.Team().GetActiveMemberCount(teamId1, nil); err != nil { t.Fatal(err) } else { if result != 1 { @@ -1138,7 +1138,7 @@ func testTeamStoreMemberCount(t *testing.T, ss store.Store) { _, err = ss.Team().SaveMember(m3, -1) require.Nil(t, err) - if totalMemberCount, err := ss.Team().GetTotalMemberCount(teamId1); err != nil { + if totalMemberCount, err := ss.Team().GetTotalMemberCount(teamId1, nil); err != nil { t.Fatal(err) } else { if totalMemberCount != 2 { @@ -1146,7 +1146,7 @@ func testTeamStoreMemberCount(t *testing.T, ss store.Store) { } } - if result, err := ss.Team().GetActiveMemberCount(teamId1); err != nil { + if result, err := ss.Team().GetActiveMemberCount(teamId1, nil); err != nil { t.Fatal(err) } else { if result != 1 {