diff --git a/api4/graphql.go b/api4/graphql.go index 1e7879290b..fa0db0df8f 100644 --- a/api4/graphql.go +++ b/api4/graphql.go @@ -65,6 +65,7 @@ const ( webCtx ctxKey = 0 rolesLoaderCtx ctxKey = 1 channelsLoaderCtx ctxKey = 2 + teamsLoaderCtx ctxKey = 3 ) const loaderBatchCapacity = 200 @@ -107,6 +108,9 @@ func (api *API) graphQL(c *Context, w http.ResponseWriter, r *http.Request) { channelsLoader := dataloader.NewBatchedLoader(graphQLChannelsLoader, dataloader.WithBatchCapacity(loaderBatchCapacity)) reqCtx = context.WithValue(reqCtx, channelsLoaderCtx, channelsLoader) + teamsLoader := dataloader.NewBatchedLoader(graphQLTeamsLoader, dataloader.WithBatchCapacity(loaderBatchCapacity)) + reqCtx = context.WithValue(reqCtx, teamsLoaderCtx, teamsLoader) + response = api.schema.Exec(reqCtx, params.Query, params.OperationName, diff --git a/api4/resolver.go b/api4/resolver.go index 08f3de8260..d0c138eec3 100644 --- a/api4/resolver.go +++ b/api4/resolver.go @@ -303,3 +303,12 @@ func getChannelsLoader(ctx context.Context) (*dataloader.Loader, error) { } return l, nil } + +// getTeamsLoader returns the teams loader out of the context. +func getTeamsLoader(ctx context.Context) (*dataloader.Loader, error) { + l, ok := ctx.Value(teamsLoaderCtx).(*dataloader.Loader) + if !ok { + return nil, errors.New("no dataloader.Loader found in context") + } + return l, nil +} diff --git a/api4/resolver_team.go b/api4/resolver_team.go index 83779f960e..07d911d47f 100644 --- a/api4/resolver_team.go +++ b/api4/resolver_team.go @@ -5,8 +5,12 @@ package api4 import ( "context" + "fmt" + + "github.com/graph-gophers/dataloader/v6" "github.com/mattermost/mattermost-server/v6/model" + "github.com/mattermost/mattermost-server/v6/web" ) func getGraphQLTeam(ctx context.Context, id string) (*model.Team, error) { @@ -15,11 +19,18 @@ func getGraphQLTeam(ctx context.Context, id string) (*model.Team, error) { return nil, err } - team, appErr := c.App.GetTeam(id) - if appErr != nil { - return nil, appErr + loader, err := getTeamsLoader(ctx) + if err != nil { + return nil, err } + thunk := loader.Load(ctx, dataloader.StringKey(id)) + result, err := thunk() + if err != nil { + return nil, err + } + team := result.(*model.Team) + if (!team.AllowOpenInvite || team.Type != model.TeamOpen) && !c.App.SessionHasPermissionToTeam(*c.AppContext.Session(), team.Id, model.PermissionViewTeam) { c.SetPermissionError(model.PermissionViewTeam) @@ -29,3 +40,53 @@ func getGraphQLTeam(ctx context.Context, id string) (*model.Team, error) { team = c.App.SanitizeTeam(*c.AppContext.Session(), team) return team, nil } + +func graphQLTeamsLoader(ctx context.Context, keys dataloader.Keys) []*dataloader.Result { + stringKeys := keys.Keys() + result := make([]*dataloader.Result, len(stringKeys)) + + c, err := getCtx(ctx) + if err != nil { + for i := range result { + result[i] = &dataloader.Result{Error: err} + } + return result + } + + teams, err := getGraphQLTeams(c, stringKeys) + if err != nil { + for i := range result { + result[i] = &dataloader.Result{Error: err} + } + return result + } + + for i, ch := range teams { + result[i] = &dataloader.Result{Data: ch} + } + return result +} + +func getGraphQLTeams(c *web.Context, teamIDs []string) ([]*model.Team, error) { + teams, appErr := c.App.GetTeams(teamIDs) + if appErr != nil { + return nil, appErr + } + + if len(teams) != len(teamIDs) { + return nil, fmt.Errorf("All teams were not found. Requested %d; Found %d", len(teamIDs), len(teams)) + } + + // The teams need to be in the exact same order as the input slice. + tmp := make(map[string]*model.Team) + for _, ch := range teams { + tmp[ch.Id] = ch + } + + // We reuse the same slice and just rewrite the teams. + for i, id := range teamIDs { + teams[i] = tmp[id] + } + + return teams, nil +} diff --git a/app/app_iface.go b/app/app_iface.go index cfc816b797..6816e9d4bd 100644 --- a/app/app_iface.go +++ b/app/app_iface.go @@ -742,6 +742,7 @@ type AppIface interface { GetTeamPoliciesForUser(userID string, offset, limit int) (*model.RetentionPolicyForTeamList, *model.AppError) GetTeamStats(teamID string, restrictions *model.ViewUsersRestrictions) (*model.TeamStats, *model.AppError) GetTeamUnread(teamID, userID string) (*model.TeamUnread, *model.AppError) + GetTeams(teamIDs []string) ([]*model.Team, *model.AppError) GetTeamsForRetentionPolicy(policyID string, offset, limit int) (*model.TeamsWithCount, *model.AppError) GetTeamsForScheme(scheme *model.Scheme, offset int, limit int) ([]*model.Team, *model.AppError) GetTeamsForSchemePage(scheme *model.Scheme, page int, perPage int) ([]*model.Team, *model.AppError) diff --git a/app/opentracing/opentracing_layer.go b/app/opentracing/opentracing_layer.go index 13a7a7cc1f..3e25ac8b1a 100644 --- a/app/opentracing/opentracing_layer.go +++ b/app/opentracing/opentracing_layer.go @@ -9332,6 +9332,28 @@ func (a *OpenTracingAppLayer) GetTeamUnread(teamID string, userID string) (*mode return resultVar0, resultVar1 } +func (a *OpenTracingAppLayer) GetTeams(teamIDs []string) ([]*model.Team, *model.AppError) { + origCtx := a.ctx + span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetTeams") + + a.ctx = newCtx + a.app.Srv().Store.SetContext(newCtx) + defer func() { + a.app.Srv().Store.SetContext(origCtx) + a.ctx = origCtx + }() + + defer span.Finish() + resultVar0, resultVar1 := a.app.GetTeams(teamIDs) + + if resultVar1 != nil { + span.LogFields(spanlog.Error(resultVar1)) + ext.Error.Set(span, true) + } + + return resultVar0, resultVar1 +} + func (a *OpenTracingAppLayer) GetTeamsForRetentionPolicy(policyID string, offset int, limit int) (*model.TeamsWithCount, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetTeamsForRetentionPolicy") diff --git a/app/team.go b/app/team.go index 6fb01863e7..b3b131ab0b 100644 --- a/app/team.go +++ b/app/team.go @@ -741,6 +741,21 @@ func (a *App) GetTeam(teamID string) (*model.Team, *model.AppError) { return team, nil } +func (a *App) GetTeams(teamIDs []string) ([]*model.Team, *model.AppError) { + teams, err := a.ch.srv.teamService.GetTeams(teamIDs) + if err != nil { + var nfErr *store.ErrNotFound + switch { + case errors.As(err, &nfErr): + return nil, model.NewAppError("GetTeam", "app.team.get.find.app_error", nil, nfErr.Error(), http.StatusNotFound) + default: + return nil, model.NewAppError("GetTeam", "app.team.get.finding.app_error", nil, err.Error(), http.StatusInternalServerError) + } + } + + return teams, nil +} + func (a *App) GetTeamByName(name string) (*model.Team, *model.AppError) { team, err := a.Srv().Store.Team().GetByName(name) if err != nil { diff --git a/app/teams/teams.go b/app/teams/teams.go index 1d7ddbc412..a143274a30 100644 --- a/app/teams/teams.go +++ b/app/teams/teams.go @@ -33,6 +33,15 @@ func (ts *TeamService) GetTeam(teamID string) (*model.Team, error) { return team, nil } +func (ts *TeamService) GetTeams(teamIDs []string) ([]*model.Team, error) { + teams, err := ts.store.GetMany(teamIDs) + if err != nil { + return nil, err + } + + return teams, nil +} + // CreateDefaultChannels creates channels in the given team for each channel returned by (*App).DefaultChannelNames. // func (ts *TeamService) createDefaultChannels(teamID string) ([]*model.Channel, error) { diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index ca33545bd0..28720afbba 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -8624,6 +8624,24 @@ func (s *OpenTracingLayerTeamStore) GetCommonTeamIDsForTwoUsers(userID string, o return result, err } +func (s *OpenTracingLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetMany") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.TeamStore.GetMany(ids) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetMember") diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index 430f08c958..10f6f33bb9 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -9841,6 +9841,27 @@ func (s *RetryLayerTeamStore) GetCommonTeamIDsForTwoUsers(userID string, otherUs } +func (s *RetryLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) { + + tries := 0 + for { + result, err := s.TeamStore.GetMany(ids) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { tries := 0 diff --git a/store/sqlstore/team_store.go b/store/sqlstore/team_store.go index b852752098..d576afeb1f 100644 --- a/store/sqlstore/team_store.go +++ b/store/sqlstore/team_store.go @@ -303,6 +303,29 @@ func (s SqlTeamStore) Get(id string) (*model.Team, error) { return &team, nil } +func (s SqlTeamStore) GetMany(ids []string) ([]*model.Team, error) { + query := s.getQueryBuilder(). + Select("*"). + From("Teams"). + Where(sq.Eq{"Id": ids}) + sql, args, err := query.ToSql() + if err != nil { + return nil, errors.Wrapf(err, "getmany_tosql") + } + + teams := []*model.Team{} + err = s.GetReplicaX().Select(&teams, sql, args...) + if err != nil { + return nil, errors.Wrapf(err, "failed to get teams with ids %v", ids) + } + + if len(teams) == 0 { + return nil, store.NewErrNotFound("Team", fmt.Sprintf("ids=%v", ids)) + } + + return teams, nil +} + // GetByInviteId returns from the database the team that matches the inviteId provided as parameter. // If the parameter provided is empty or if there is no match in the database, it returns a model.AppError // with a http.StatusNotFound in the StatusCode field. diff --git a/store/store.go b/store/store.go index 83a37ddfc5..18ceb7f7bc 100644 --- a/store/store.go +++ b/store/store.go @@ -102,6 +102,7 @@ type TeamStore interface { Save(team *model.Team) (*model.Team, error) Update(team *model.Team) (*model.Team, error) Get(id string) (*model.Team, error) + GetMany(ids []string) ([]*model.Team, error) GetByName(name string) (*model.Team, error) GetByNames(name []string) ([]*model.Team, error) SearchAll(opts *model.TeamSearch) ([]*model.Team, error) diff --git a/store/storetest/mocks/TeamStore.go b/store/storetest/mocks/TeamStore.go index 746937bcf9..388c579dde 100644 --- a/store/storetest/mocks/TeamStore.go +++ b/store/storetest/mocks/TeamStore.go @@ -374,6 +374,29 @@ func (_m *TeamStore) GetCommonTeamIDsForTwoUsers(userID string, otherUserID stri return r0, r1 } +// GetMany provides a mock function with given fields: ids +func (_m *TeamStore) GetMany(ids []string) ([]*model.Team, error) { + ret := _m.Called(ids) + + var r0 []*model.Team + if rf, ok := ret.Get(0).(func([]string) []*model.Team); ok { + r0 = rf(ids) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.Team) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func([]string) error); ok { + r1 = rf(ids) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // GetMember provides a mock function with given fields: ctx, teamID, userID func (_m *TeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { ret := _m.Called(ctx, teamID, userID) diff --git a/store/storetest/team_store.go b/store/storetest/team_store.go index cbad0cb7fd..2104aa41c8 100644 --- a/store/storetest/team_store.go +++ b/store/storetest/team_store.go @@ -31,6 +31,7 @@ func TestTeamStore(t *testing.T, ss store.Store) { t.Run("Save", func(t *testing.T) { testTeamStoreSave(t, ss) }) t.Run("Update", func(t *testing.T) { testTeamStoreUpdate(t, ss) }) t.Run("Get", func(t *testing.T) { testTeamStoreGet(t, ss) }) + t.Run("GetMany", func(t *testing.T) { testTeamStoreGetMany(t, ss) }) t.Run("GetByName", func(t *testing.T) { testTeamStoreGetByName(t, ss) }) t.Run("GetByNames", func(t *testing.T) { testTeamStoreGetByNames(t, ss) }) t.Run("SearchAll", func(t *testing.T) { testTeamStoreSearchAll(t, ss) }) @@ -132,6 +133,37 @@ func testTeamStoreGet(t *testing.T, ss store.Store) { require.Error(t, err, "Missing id should have failed") } +func testTeamStoreGetMany(t *testing.T, ss store.Store) { + o1, err := ss.Team().Save(&model.Team{ + DisplayName: "DisplayName", + Name: NewTestId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + o2, err := ss.Team().Save(&model.Team{ + DisplayName: "DisplayName2", + Name: NewTestId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + res, err := ss.Team().GetMany([]string{o1.Id, o2.Id}) + require.NoError(t, err) + assert.Len(t, res, 2) + + res, err = ss.Team().GetMany([]string{o1.Id, "notexists"}) + require.NoError(t, err) + assert.Len(t, res, 1) + + _, err = ss.Team().GetMany([]string{"whereisit", "notexists"}) + require.Error(t, err) + var nfErr *store.ErrNotFound + assert.True(t, errors.As(err, &nfErr)) +} + func testTeamStoreGetByNames(t *testing.T, ss store.Store) { o1 := model.Team{} o1.DisplayName = "DisplayName" diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index 85807d2f09..e3fe87c9ca 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -7768,6 +7768,22 @@ func (s *TimerLayerTeamStore) GetCommonTeamIDsForTwoUsers(userID string, otherUs return result, err } +func (s *TimerLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) { + start := timemodule.Now() + + result, err := s.TeamStore.GetMany(ids) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("TeamStore.GetMany", success, elapsed) + } + return result, err +} + func (s *TimerLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { start := timemodule.Now()