diff --git a/server/channels/api4/apitestlib.go b/server/channels/api4/apitestlib.go index 963471074a..449d4c6208 100644 --- a/server/channels/api4/apitestlib.go +++ b/server/channels/api4/apitestlib.go @@ -1175,7 +1175,7 @@ func (th *TestHelper) MakeUserChannelAdmin(user *model.User, channel *model.Chan } func (th *TestHelper) UpdateUserToTeamAdmin(user *model.User, team *model.Team) { - if tm, err := th.App.Srv().Store().Team().GetMember(context.Background(), team.Id, user.Id); err == nil { + if tm, err := th.App.Srv().Store().Team().GetMember(th.Context, team.Id, user.Id); err == nil { tm.SchemeAdmin = true if _, err = th.App.Srv().Store().Team().UpdateMember(tm); err != nil { panic(err) @@ -1186,7 +1186,7 @@ func (th *TestHelper) UpdateUserToTeamAdmin(user *model.User, team *model.Team) } func (th *TestHelper) UpdateUserToNonTeamAdmin(user *model.User, team *model.Team) { - if tm, err := th.App.Srv().Store().Team().GetMember(context.Background(), team.Id, user.Id); err == nil { + if tm, err := th.App.Srv().Store().Team().GetMember(th.Context, team.Id, user.Id); err == nil { tm.SchemeAdmin = false if _, err = th.App.Srv().Store().Team().UpdateMember(tm); err != nil { panic(err) diff --git a/server/channels/api4/channel.go b/server/channels/api4/channel.go index b703196a8a..03a37c2327 100644 --- a/server/channels/api4/channel.go +++ b/server/channels/api4/channel.go @@ -465,7 +465,7 @@ func createDirectChannel(c *Context, w http.ResponseWriter, r *http.Request) { audit.AddEventParameter(auditRec, "user_id", otherUserId) - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, otherUserId) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, otherUserId) if err != nil { c.Err = err return @@ -546,7 +546,7 @@ func createGroupChannel(c *Context, w http.ResponseWriter, r *http.Request) { canSeeAll := true for _, id := range userIds { if c.AppContext.Session().UserId != id { - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, id) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, id) if err != nil { c.Err = err return @@ -1110,7 +1110,7 @@ func searchChannelsForTeam(c *Context, w http.ResponseWriter, r *http.Request) { channels, appErr = c.App.SearchChannels(c.AppContext, c.Params.TeamId, props.Term) } else { // If the user is not a team member, return a 404 - if _, appErr = c.App.GetTeamMember(c.Params.TeamId, c.AppContext.Session().UserId); appErr != nil { + if _, appErr = c.App.GetTeamMember(c.AppContext, c.Params.TeamId, c.AppContext.Session().UserId); appErr != nil { c.Err = appErr return } @@ -1149,7 +1149,7 @@ func searchArchivedChannelsForTeam(c *Context, w http.ResponseWriter, r *http.Re channels, appErr = c.App.SearchArchivedChannels(c.AppContext, c.Params.TeamId, props.Term, c.AppContext.Session().UserId) } else { // If the user is not a team member, return a 404 - if _, appErr = c.App.GetTeamMember(c.Params.TeamId, c.AppContext.Session().UserId); appErr != nil { + if _, appErr = c.App.GetTeamMember(c.AppContext, c.Params.TeamId, c.AppContext.Session().UserId); appErr != nil { c.Err = appErr return } diff --git a/server/channels/api4/emoji.go b/server/channels/api4/emoji.go index d6a08d2aec..3f9ed7dedf 100644 --- a/server/channels/api4/emoji.go +++ b/server/channels/api4/emoji.go @@ -52,7 +52,7 @@ func createEmoji(c *Context, w http.ResponseWriter, r *http.Request) { defer c.LogAuditRec(auditRec) // Allow any user with CREATE_EMOJIS permission at Team level to create emojis at system level - memberships, err := c.App.GetTeamMembersForUser(c.AppContext.Session().UserId, "", true) + memberships, err := c.App.GetTeamMembersForUser(c.AppContext, c.AppContext.Session().UserId, "", true) if err != nil { c.Err = err @@ -144,7 +144,7 @@ func deleteEmoji(c *Context, w http.ResponseWriter, r *http.Request) { auditRec.AddEventObjectType("emoji") // Allow any user with DELETE_EMOJIS permission at Team level to delete emojis at system level - memberships, err := c.App.GetTeamMembersForUser(c.AppContext.Session().UserId, "", true) + memberships, err := c.App.GetTeamMembersForUser(c.AppContext, c.AppContext.Session().UserId, "", true) if err != nil { c.Err = err diff --git a/server/channels/api4/group.go b/server/channels/api4/group.go index 4be9d5989c..d053f18a56 100644 --- a/server/channels/api4/group.go +++ b/server/channels/api4/group.go @@ -109,7 +109,7 @@ func getGroup(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -686,7 +686,7 @@ func getGroupMembers(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -1061,7 +1061,7 @@ func getGroups(c *Context, w http.ResponseWriter, r *http.Request) { opts.Since = since } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -1073,7 +1073,7 @@ func getGroups(c *Context, w http.ResponseWriter, r *http.Request) { ) if opts.FilterHasMember != "" { - canSee, appErr = c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, opts.FilterHasMember) + canSee, appErr = c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, opts.FilterHasMember) if appErr != nil { c.Err = appErr return diff --git a/server/channels/api4/post_test.go b/server/channels/api4/post_test.go index 3c6287ab21..0d28c057ac 100644 --- a/server/channels/api4/post_test.go +++ b/server/channels/api4/post_test.go @@ -425,7 +425,7 @@ func TestCreatePostWithOAuthClient(t *testing.T) { }) require.Nil(t, appErr, "should create an OAuthApp") - session, appErr := th.App.CreateSession(&model.Session{ + session, appErr := th.App.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, Token: "token", IsOAuth: true, @@ -763,7 +763,7 @@ func TestCreatePostPublic(t *testing.T) { th.App.UpdateUserRoles(th.Context, ruser.Id, model.SystemUserRoleId, false) th.App.JoinUserToTeam(th.Context, th.BasicTeam, ruser, "") - th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, model.TeamUserRoleId+" "+model.TeamPostAllPublicRoleId) + th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, model.TeamUserRoleId+" "+model.TeamPostAllPublicRoleId) th.App.Srv().InvalidateAllCaches() client.Login(context.Background(), user.Email, user.Password) @@ -816,7 +816,7 @@ func TestCreatePostAll(t *testing.T) { th.App.UpdateUserRoles(th.Context, ruser.Id, model.SystemUserRoleId, false) th.App.JoinUserToTeam(th.Context, th.BasicTeam, ruser, "") - th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, model.TeamUserRoleId+" "+model.TeamPostAllRoleId) + th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, model.TeamUserRoleId+" "+model.TeamPostAllRoleId) th.App.Srv().InvalidateAllCaches() client.Login(context.Background(), user.Email, user.Password) diff --git a/server/channels/api4/resolver.go b/server/channels/api4/resolver.go index ce100f4647..bdd1b79329 100644 --- a/server/channels/api4/resolver.go +++ b/server/channels/api4/resolver.go @@ -151,7 +151,7 @@ func (r *resolver) TeamMembers(ctx context.Context, args struct { return nil, c.Err } - canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, args.UserID) + canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, args.UserID) if appErr != nil { return nil, appErr } @@ -167,7 +167,7 @@ func (r *resolver) TeamMembers(ctx context.Context, args struct { return nil, c.Err } - tm, appErr2 := c.App.GetTeamMember(args.TeamID, args.UserID) + tm, appErr2 := c.App.GetTeamMember(c.AppContext, args.TeamID, args.UserID) if appErr2 != nil { return nil, appErr2 } @@ -181,7 +181,7 @@ func (r *resolver) TeamMembers(ctx context.Context, args struct { } // Do not return archived team members - members, appErr := c.App.GetTeamMembersForUser(args.UserID, excludeTeamID, false) + members, appErr := c.App.GetTeamMembersForUser(c.AppContext, args.UserID, excludeTeamID, false) if appErr != nil { return nil, appErr } diff --git a/server/channels/api4/resolver_user.go b/server/channels/api4/resolver_user.go index 868e2555f1..dcacd6f61c 100644 --- a/server/channels/api4/resolver_user.go +++ b/server/channels/api4/resolver_user.go @@ -139,7 +139,7 @@ func (u *user) Sessions(ctx context.Context) ([]*model.Session, error) { return nil, c.Err } - sessions, appErr := c.App.GetSessions(u.Id) + sessions, appErr := c.App.GetSessions(c.AppContext, u.Id) if appErr != nil { return nil, appErr } @@ -183,7 +183,7 @@ func getGraphQLUsers(c *web.Context, userIDs []string) ([]*model.User, error) { // and cached for the rest of the query. So it's not an issue // to run this in a loop. for _, id := range userIDs { - canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, id) + canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, id) if appErr != nil || !canSee { c.SetPermissionError(model.PermissionViewMembers) return nil, c.Err diff --git a/server/channels/api4/team.go b/server/channels/api4/team.go index 24349b8806..8c6ac8e9ec 100644 --- a/server/channels/api4/team.go +++ b/server/channels/api4/team.go @@ -537,7 +537,7 @@ func getTeamMember(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, c.Params.UserId) + canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, c.Params.UserId) if appErr != nil { c.Err = appErr return @@ -548,7 +548,7 @@ func getTeamMember(c *Context, w http.ResponseWriter, r *http.Request) { return } - team, appErr := c.App.GetTeamMember(c.Params.TeamId, c.Params.UserId) + team, appErr := c.App.GetTeamMember(c.AppContext, c.Params.TeamId, c.Params.UserId) if appErr != nil { c.Err = appErr return @@ -574,7 +574,7 @@ func getTeamMembers(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -612,7 +612,7 @@ func getTeamMembersForUser(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, c.Params.UserId) + canSee, appErr := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, c.Params.UserId) if appErr != nil { c.Err = appErr return @@ -623,7 +623,7 @@ func getTeamMembersForUser(c *Context, w http.ResponseWriter, r *http.Request) { return } - members, appErr := c.App.GetTeamMembersForUser(c.Params.UserId, "", true) + members, appErr := c.App.GetTeamMembersForUser(c.AppContext, c.Params.UserId, "", true) if appErr != nil { c.Err = appErr return @@ -656,7 +656,7 @@ func getTeamMembersByIds(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -1007,7 +1007,7 @@ func getTeamStats(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, err := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, err := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if err != nil { c.Err = err return @@ -1047,7 +1047,7 @@ func updateTeamMemberRoles(c *Context, w http.ResponseWriter, r *http.Request) { return } - teamMember, err := c.App.UpdateTeamMemberRoles(c.Params.TeamId, c.Params.UserId, newRoles) + teamMember, err := c.App.UpdateTeamMemberRoles(c.AppContext, c.Params.TeamId, c.Params.UserId, newRoles) if err != nil { c.Err = err return @@ -1081,7 +1081,7 @@ func updateTeamMemberSchemeRoles(c *Context, w http.ResponseWriter, r *http.Requ return } - teamMember, err := c.App.UpdateTeamMemberSchemeRoles(c.Params.TeamId, c.Params.UserId, schemeRoles.SchemeGuest, schemeRoles.SchemeUser, schemeRoles.SchemeAdmin) + teamMember, err := c.App.UpdateTeamMemberSchemeRoles(c.AppContext, c.Params.TeamId, c.Params.UserId, schemeRoles.SchemeGuest, schemeRoles.SchemeUser, schemeRoles.SchemeAdmin) if err != nil { c.Err = err return @@ -1235,7 +1235,7 @@ func teamExists(c *Context, w http.ResponseWriter, r *http.Request) { if team != nil { var teamMember *model.TeamMember - teamMember, err = c.App.GetTeamMember(team.Id, c.AppContext.Session().UserId) + teamMember, err = c.App.GetTeamMember(c.AppContext, team.Id, c.AppContext.Session().UserId) if err != nil && err.StatusCode != http.StatusNotFound { c.Err = err return diff --git a/server/channels/api4/user.go b/server/channels/api4/user.go index 900a85b7de..befdab78a0 100644 --- a/server/channels/api4/user.go +++ b/server/channels/api4/user.go @@ -182,7 +182,7 @@ func getUser(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, c.Params.UserId) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, c.Params.UserId) if err != nil { c.SetPermissionError(model.PermissionViewMembers) return @@ -238,7 +238,7 @@ func getUserByUsername(c *Context, w http.ResponseWriter, r *http.Request) { user, err := c.App.GetUserByUsername(c.Params.Username) if err != nil { - restrictions, err2 := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, err2 := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if err2 != nil { c.Err = err2 return @@ -251,7 +251,7 @@ func getUserByUsername(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, user.Id) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, user.Id) if err != nil { c.Err = err return @@ -306,7 +306,7 @@ func getUserByEmail(c *Context, w http.ResponseWriter, r *http.Request) { user, err := c.App.GetUserByEmail(c.Params.Email) if err != nil { - restrictions, err2 := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, err2 := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if err2 != nil { c.Err = err2 return @@ -319,7 +319,7 @@ func getUserByEmail(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, user.Id) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, user.Id) if err != nil { c.Err = err return @@ -349,7 +349,7 @@ func getDefaultProfileImage(c *Context, w http.ResponseWriter, r *http.Request) return } - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, c.Params.UserId) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, c.Params.UserId) if err != nil { c.Err = err return @@ -383,7 +383,7 @@ func getProfileImage(c *Context, w http.ResponseWriter, r *http.Request) { return } - canSee, err := c.App.UserCanSeeOtherUser(c.AppContext.Session().UserId, c.Params.UserId) + canSee, err := c.App.UserCanSeeOtherUser(c.AppContext, c.AppContext.Session().UserId, c.Params.UserId) if err != nil { c.Err = err return @@ -538,7 +538,7 @@ func getTotalUsersStats(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, err := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, err := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if err != nil { c.Err = err return @@ -769,7 +769,7 @@ func getUsers(c *Context, w http.ResponseWriter, r *http.Request) { } } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -906,7 +906,7 @@ func getUsers(c *Context, w http.ResponseWriter, r *http.Request) { return } } else { - userGetOptions, appErr = c.App.RestrictUsersGetByPermissions(c.AppContext.Session().UserId, userGetOptions) + userGetOptions, appErr = c.App.RestrictUsersGetByPermissions(c.AppContext, c.AppContext.Session().UserId, userGetOptions) if appErr != nil { c.Err = appErr return @@ -979,7 +979,7 @@ func getUsersByIds(c *Context, w http.ResponseWriter, r *http.Request) { options.Since = since } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -1009,7 +1009,7 @@ func getUsersByNames(c *Context, w http.ResponseWriter, r *http.Request) { return } - restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext.Session().UserId) + restrictions, appErr := c.App.GetViewUsersRestrictions(c.AppContext, c.AppContext.Session().UserId) if appErr != nil { c.Err = appErr return @@ -1124,7 +1124,7 @@ func searchUsers(c *Context, w http.ResponseWriter, r *http.Request) { options.AllowFullNames = *c.App.Config().PrivacySettings.ShowFullName } - options, appErr := c.App.RestrictUsersSearchByPermissions(c.AppContext.Session().UserId, options) + options, appErr := c.App.RestrictUsersSearchByPermissions(c.AppContext, c.AppContext.Session().UserId, options) if appErr != nil { c.Err = appErr return @@ -1187,7 +1187,7 @@ func autocompleteUsers(c *Context, w http.ResponseWriter, r *http.Request) { var autocomplete model.UserAutocomplete var err *model.AppError - options, err = c.App.RestrictUsersSearchByPermissions(c.AppContext.Session().UserId, options) + options, err = c.App.RestrictUsersSearchByPermissions(c.AppContext, c.AppContext.Session().UserId, options) if err != nil { c.Err = err return @@ -2081,7 +2081,7 @@ func Logout(c *Context, w http.ResponseWriter, r *http.Request) { c.RemoveSessionCookie(w, r) if c.AppContext.Session().Id != "" { - if err := c.App.RevokeSessionById(c.AppContext.Session().Id); err != nil { + if err := c.App.RevokeSessionById(c.AppContext, c.AppContext.Session().Id); err != nil { c.Err = err return } @@ -2102,7 +2102,7 @@ func getSessions(c *Context, w http.ResponseWriter, r *http.Request) { return } - sessions, appErr := c.App.GetSessions(c.Params.UserId) + sessions, appErr := c.App.GetSessions(c.AppContext, c.Params.UserId) if appErr != nil { c.Err = appErr return @@ -2143,7 +2143,7 @@ func revokeSession(c *Context, w http.ResponseWriter, r *http.Request) { } audit.AddEventParameter(auditRec, "session_id", sessionId) - session, err := c.App.GetSessionById(sessionId) + session, err := c.App.GetSessionById(c.AppContext, sessionId) if err != nil { c.Err = err return @@ -2157,7 +2157,7 @@ func revokeSession(c *Context, w http.ResponseWriter, r *http.Request) { return } - if err := c.App.RevokeSession(session); err != nil { + if err := c.App.RevokeSession(c.AppContext, session); err != nil { c.Err = err return } @@ -2183,7 +2183,7 @@ func revokeAllSessionsForUser(c *Context, w http.ResponseWriter, r *http.Request return } - if err := c.App.RevokeAllSessions(c.Params.UserId); err != nil { + if err := c.App.RevokeAllSessions(c.AppContext, c.Params.UserId); err != nil { c.Err = err return } @@ -2228,7 +2228,7 @@ func attachDeviceId(c *Context, w http.ResponseWriter, r *http.Request) { audit.AddEventParameter(auditRec, "device_id", deviceId) // A special case where we logout of all other sessions with the same device id - if err := c.App.RevokeSessionsForDeviceId(c.AppContext.Session().UserId, deviceId, c.AppContext.Session().Id); err != nil { + if err := c.App.RevokeSessionsForDeviceId(c.AppContext, c.AppContext.Session().UserId, deviceId, c.AppContext.Session().Id); err != nil { c.Err = err return } @@ -2388,7 +2388,7 @@ func switchAccountType(c *Context, w http.ResponseWriter, r *http.Request) { return } - link, err = c.App.SwitchOAuthToEmail(switchRequest.Email, switchRequest.NewPassword, c.AppContext.Session().UserId) + link, err = c.App.SwitchOAuthToEmail(c.AppContext, switchRequest.Email, switchRequest.NewPassword, c.AppContext.Session().UserId) } else if switchRequest.EmailToLdap() { link, err = c.App.SwitchEmailToLdap(c.AppContext, switchRequest.Email, switchRequest.Password, switchRequest.MfaCode, switchRequest.LdapLoginId, switchRequest.NewPassword) } else if switchRequest.LdapToEmail() { @@ -2623,7 +2623,7 @@ func revokeUserAccessToken(c *Context, w http.ResponseWriter, r *http.Request) { return } - if err = c.App.RevokeUserAccessToken(accessToken); err != nil { + if err = c.App.RevokeUserAccessToken(c.AppContext, accessToken); err != nil { c.Err = err return } @@ -2668,7 +2668,7 @@ func disableUserAccessToken(c *Context, w http.ResponseWriter, r *http.Request) return } - if err = c.App.DisableUserAccessToken(accessToken); err != nil { + if err = c.App.DisableUserAccessToken(c.AppContext, accessToken); err != nil { c.Err = err return } @@ -2713,7 +2713,7 @@ func enableUserAccessToken(c *Context, w http.ResponseWriter, r *http.Request) { return } - if err = c.App.EnableUserAccessToken(accessToken); err != nil { + if err = c.App.EnableUserAccessToken(c.AppContext, accessToken); err != nil { c.Err = err return } diff --git a/server/channels/api4/user_test.go b/server/channels/api4/user_test.go index a5917631c7..9ab83a0756 100644 --- a/server/channels/api4/user_test.go +++ b/server/channels/api4/user_test.go @@ -3475,7 +3475,7 @@ func TestRevokeSessions(t *testing.T) { th.LoginBasic() - sessions, _ = th.App.GetSessions(th.SystemAdminUser.Id) + sessions, _ = th.App.GetSessions(th.Context, th.SystemAdminUser.Id) session = sessions[0] resp, err = th.Client.RevokeSession(context.Background(), user.Id, session.Id) @@ -3559,10 +3559,10 @@ func TestRevokeSessionsFromAllUsers(t *testing.T) { th.Client.Login(context.Background(), user.Email, user.Password) admin := th.SystemAdminUser th.Client.Login(context.Background(), admin.Email, admin.Password) - sessions, err := th.Server.Store().Session().GetSessions(user.Id) + sessions, err := th.Server.Store().Session().GetSessions(th.Context, user.Id) require.NotEmpty(t, sessions) require.NoError(t, err) - sessions, err = th.Server.Store().Session().GetSessions(admin.Id) + sessions, err = th.Server.Store().Session().GetSessions(th.Context, admin.Id) require.NotEmpty(t, sessions) require.NoError(t, err) _, err = th.Client.RevokeSessionsFromAllUsers(context.Background()) @@ -3574,11 +3574,11 @@ func TestRevokeSessionsFromAllUsers(t *testing.T) { require.Error(t, err) CheckUnauthorizedStatus(t, resp) - sessions, err = th.Server.Store().Session().GetSessions(user.Id) + sessions, err = th.Server.Store().Session().GetSessions(th.Context, user.Id) require.Empty(t, sessions) require.NoError(t, err) - sessions, err = th.Server.Store().Session().GetSessions(admin.Id) + sessions, err = th.Server.Store().Session().GetSessions(th.Context, admin.Id) require.Empty(t, sessions) require.NoError(t, err) } @@ -3611,7 +3611,7 @@ func TestAttachDeviceId(t *testing.T) { cookies := resp.Header.Get("Set-Cookie") assert.Regexp(t, tc.ExpectedSetCookieHeaderRegexp, cookies) - sessions, appErr := th.App.GetSessions(th.BasicUser.Id) + sessions, appErr := th.App.GetSessions(th.Context, th.BasicUser.Id) require.Nil(t, appErr) assert.Equal(t, deviceId, sessions[0].DeviceId, "Missing device Id") }) @@ -3889,7 +3889,7 @@ func TestLoginWithLag(t *testing.T) { mainHelper.SQLStore.UpdateLicense(model.NewTestLicense("ldap")) mainHelper.ToggleReplicasOff() - appErr := th.App.RevokeAllSessions(th.BasicUser.Id) + appErr := th.App.RevokeAllSessions(th.Context, th.BasicUser.Id) require.Nil(t, appErr) mainHelper.ToggleReplicasOn() diff --git a/server/channels/app/app_iface.go b/server/channels/app/app_iface.go index 9a822ab8cf..e1be1c577e 100644 --- a/server/channels/app/app_iface.go +++ b/server/channels/app/app_iface.go @@ -39,7 +39,7 @@ import ( // AppIface is extracted from App struct and contains all it's exported methods. It's provided to allow partial interface passing and app layers creation. type AppIface interface { // @openTracingParams args - ExecuteCommand(c request.CTX, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) + ExecuteCommand(c *request.Context, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) // @openTracingParams teamID // previous ListCommands now ListAutocompleteCommands ListAutocompleteCommands(teamID string, T i18n.TranslateFunc) ([]*model.Command, *model.AppError) @@ -129,7 +129,7 @@ type AppIface interface { DeletePublicKey(name string) *model.AppError // DemoteUserToGuest Convert user's roles and all his membership's roles from // regular user roles to guest roles. - DemoteUserToGuest(c request.CTX, user *model.User) *model.AppError + DemoteUserToGuest(c *request.Context, user *model.User) *model.AppError // DisablePlugin will set the config for an installed plugin to disabled, triggering deactivation if active. // Notifies cluster peers through config change. DisablePlugin(id string) *model.AppError @@ -365,7 +365,7 @@ type AppIface interface { // This to be used for places we check the users password when they are already logged in DoubleCheckPassword(user *model.User, password string) *model.AppError // UpdateBotActive marks a bot as active or inactive, along with its corresponding user. - UpdateBotActive(c request.CTX, botUserId string, active bool) (*model.Bot, *model.AppError) + UpdateBotActive(c *request.Context, botUserId string, active bool) (*model.Bot, *model.AppError) // UpdateBotOwner changes a bot's owner to the given value. UpdateBotOwner(botUserId, newOwnerId string) (*model.Bot, *model.AppError) // UpdateChannel updates a given channel by its Id. It also publishes the CHANNEL_UPDATED event. @@ -424,7 +424,7 @@ type AppIface interface { AdjustImage(file io.Reader) (*bytes.Buffer, *model.AppError) AdjustInProductLimits(limits *model.ProductLimits, subscription *model.Subscription) *model.AppError AdjustTeamsFromProductLimits(teamLimits *model.TeamsLimits) *model.AppError - AllowOAuthAppAccessToUser(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) + AllowOAuthAppAccessToUser(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) AppendFile(fr io.Reader, path string) (int64, *model.AppError) AsymmetricSigningKey() *ecdsa.PrivateKey AttachCloudSessionCookie(c *request.Context, w http.ResponseWriter, r *http.Request) @@ -500,7 +500,7 @@ type AppIface interface { CreateRetentionPolicy(policy *model.RetentionPolicyWithTeamAndChannelIDs) (*model.RetentionPolicyWithTeamAndChannelCounts, *model.AppError) CreateRole(role *model.Role) (*model.Role, *model.AppError) CreateScheme(scheme *model.Scheme) (*model.Scheme, *model.AppError) - CreateSession(session *model.Session) (*model.Session, *model.AppError) + CreateSession(c *request.Context, session *model.Session) (*model.Session, *model.AppError) CreateSidebarCategory(c request.CTX, userID, teamID string, newCategory *model.SidebarCategoryWithChannels) (*model.SidebarCategoryWithChannels, *model.AppError) CreateTeam(c request.CTX, team *model.Team) (*model.Team, *model.AppError) CreateTeamWithUser(c *request.Context, team *model.Team, userID string) (*model.Team, *model.AppError) @@ -517,7 +517,7 @@ type AppIface interface { DataRetention() einterfaces.DataRetentionInterface DeactivateGuests(c *request.Context) *model.AppError DeactivateMfa(userID string) *model.AppError - DeauthorizeOAuthAppForUser(userID, appID string) *model.AppError + DeauthorizeOAuthAppForUser(c *request.Context, userID, appID string) *model.AppError DeleteAcknowledgementForPost(c *request.Context, postID, userID string) *model.AppError DeleteAllExpiredPluginKeys() *model.AppError DeleteAllKeysForPlugin(pluginID string) *model.AppError @@ -547,7 +547,7 @@ type AppIface interface { DeleteSidebarCategory(c request.CTX, userID, teamID, categoryId string) *model.AppError DeleteToken(token *model.Token) *model.AppError DisableAutoResponder(c request.CTX, userID string, asAdmin bool) *model.AppError - DisableUserAccessToken(token *model.UserAccessToken) *model.AppError + DisableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError DoAppMigrations() DoCheckForAdminNotifications(trial bool) *model.AppError DoCommandRequest(cmd *model.Command, p url.Values) (*model.Command, *model.CommandResponse, *model.AppError) @@ -561,7 +561,7 @@ type AppIface interface { DoUploadFile(c request.CTX, now time.Time, rawTeamId string, rawChannelId string, rawUserId string, rawFilename string, data []byte) (*model.FileInfo, *model.AppError) DoUploadFileExpectModification(c request.CTX, now time.Time, rawTeamId string, rawChannelId string, rawUserId string, rawFilename string, data []byte) (*model.FileInfo, []byte, *model.AppError) DownloadFromURL(downloadURL string) ([]byte, error) - EnableUserAccessToken(token *model.UserAccessToken) *model.AppError + EnableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError EnvironmentConfig(filter func(reflect.StructField) bool) map[string]any ExportFileBackend() filestore.FileBackend ExportFileExists(path string) (bool, *model.AppError) @@ -575,7 +575,7 @@ type AppIface interface { FileSize(path string) (int64, *model.AppError) FillInChannelProps(c request.CTX, channel *model.Channel) *model.AppError FillInChannelsProps(c request.CTX, channelList model.ChannelList) *model.AppError - FilterUsersByVisible(viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) + FilterUsersByVisible(c request.CTX, viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) FindTeamByName(name string) bool FinishSendAdminNotifyPost(trial bool, now int64, pluginBasedData map[string][]*model.NotifyAdminData) GenerateAndSaveDesktopToken(createAt int64, user *model.User) (*string, *model.AppError) @@ -696,13 +696,13 @@ type AppIface interface { GetNextPostIdFromPostList(postList *model.PostList, collapsedThreads bool) string GetNotificationNameFormat(user *model.User) string GetNumberOfChannelsOnTeam(c request.CTX, teamID string) (int, *model.AppError) - GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, code, secret, refreshToken string) (*model.AccessResponse, *model.AppError) - GetOAuthAccessTokenForImplicitFlow(userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) + GetOAuthAccessTokenForCodeFlow(c *request.Context, clientId, grantType, redirectURI, code, secret, refreshToken string) (*model.AccessResponse, *model.AppError) + GetOAuthAccessTokenForImplicitFlow(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) GetOAuthApp(appID string) (*model.OAuthApp, *model.AppError) GetOAuthApps(page, perPage int) ([]*model.OAuthApp, *model.AppError) GetOAuthAppsByCreator(userID string, page, perPage int) ([]*model.OAuthApp, *model.AppError) GetOAuthCodeRedirect(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) - GetOAuthImplicitRedirect(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) + GetOAuthImplicitRedirect(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) GetOAuthLoginEndpoint(c *request.Context, w http.ResponseWriter, r *http.Request, service, teamID, action, redirectTo, loginHint string, isMobile bool, desktopToken string) (string, *model.AppError) GetOAuthSignupEndpoint(c *request.Context, w http.ResponseWriter, r *http.Request, service, teamID string, desktopToken string) (string, *model.AppError) GetOAuthStateToken(token string) (*model.Token, *model.AppError) @@ -768,8 +768,8 @@ type AppIface interface { GetSchemes(scope string, offset int, limit int) ([]*model.Scheme, *model.AppError) GetSchemesPage(scope string, page int, perPage int) ([]*model.Scheme, *model.AppError) GetSession(token string) (*model.Session, *model.AppError) - GetSessionById(sessionID string) (*model.Session, *model.AppError) - GetSessions(userID string) ([]*model.Session, *model.AppError) + GetSessionById(c *request.Context, sessionID string) (*model.Session, *model.AppError) + GetSessions(c *request.Context, userID string) ([]*model.Session, *model.AppError) GetSharedChannel(channelID string) (*model.SharedChannel, error) GetSharedChannelRemote(id string) (*model.SharedChannelRemote, error) GetSharedChannelRemoteByIds(channelID string, remoteID string) (*model.SharedChannelRemote, error) @@ -791,10 +791,10 @@ type AppIface interface { GetTeamByName(name string) (*model.Team, *model.AppError) GetTeamIcon(team *model.Team) ([]byte, *model.AppError) GetTeamIdFromQuery(query url.Values) (string, *model.AppError) - GetTeamMember(teamID, userID string) (*model.TeamMember, *model.AppError) + GetTeamMember(c request.CTX, teamID, userID string) (*model.TeamMember, *model.AppError) GetTeamMembers(teamID string, offset int, limit int, teamMembersGetOptions *model.TeamMembersGetOptions) ([]*model.TeamMember, *model.AppError) GetTeamMembersByIds(teamID string, userIDs []string, restrictions *model.ViewUsersRestrictions) ([]*model.TeamMember, *model.AppError) - GetTeamMembersForUser(userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) + GetTeamMembersForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) GetTeamMembersForUserWithPagination(userID string, page, perPage int) ([]*model.TeamMember, *model.AppError) GetTeamPoliciesForUser(userID string, offset, limit int) (*model.RetentionPolicyForTeamList, *model.AppError) GetTeamStats(teamID string, restrictions *model.ViewUsersRestrictions) (*model.TeamStats, *model.AppError) @@ -853,7 +853,7 @@ type AppIface interface { GetUsersWithoutTeam(options *model.UserGetOptions) ([]*model.User, *model.AppError) GetUsersWithoutTeamPage(options *model.UserGetOptions, asAdmin bool) ([]*model.User, *model.AppError) GetVerifyEmailToken(token string) (*model.Token, *model.AppError) - GetViewUsersRestrictions(userID string) (*model.ViewUsersRestrictions, *model.AppError) + GetViewUsersRestrictions(c request.CTX, userID string) (*model.ViewUsersRestrictions, *model.AppError) GetWarnMetricsBot() (*model.Bot, *model.AppError) GetWarnMetricsStatus() (map[string]*model.WarnMetricStatus, *model.AppError) HTTPService() httpservice.HTTPService @@ -866,9 +866,9 @@ type AppIface interface { HandleMessageExportConfig(cfg *model.Config, appCfg *model.Config) HasPermissionTo(askingUserId string, permission *model.Permission) bool HasPermissionToChannel(c request.CTX, askingUserId string, channelID string, permission *model.Permission) bool - HasPermissionToChannelByPost(askingUserId string, postID string, permission *model.Permission) bool + HasPermissionToChannelByPost(c request.CTX, askingUserId string, postID string, permission *model.Permission) bool HasPermissionToReadChannel(c request.CTX, userID string, channel *model.Channel) bool - HasPermissionToTeam(askingUserId string, teamID string, permission *model.Permission) bool + HasPermissionToTeam(c request.CTX, askingUserId string, teamID string, permission *model.Permission) bool HasPermissionToUser(askingUserId string, userID string) bool HasSharedChannel(channelID string) (bool, error) HooksManager() *product.HooksManager @@ -992,15 +992,15 @@ type AppIface interface { RestoreChannel(c request.CTX, channel *model.Channel, userID string) (*model.Channel, *model.AppError) RestoreGroup(groupID string) (*model.Group, *model.AppError) RestoreTeam(teamID string) *model.AppError - RestrictUsersGetByPermissions(userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) - RestrictUsersSearchByPermissions(userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) + RestrictUsersGetByPermissions(c request.CTX, userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) + RestrictUsersSearchByPermissions(c request.CTX, userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) ReturnSessionToPool(session *model.Session) - RevokeAccessToken(token string) *model.AppError - RevokeAllSessions(userID string) *model.AppError - RevokeSession(session *model.Session) *model.AppError - RevokeSessionById(sessionID string) *model.AppError - RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) *model.AppError - RevokeUserAccessToken(token *model.UserAccessToken) *model.AppError + RevokeAccessToken(c *request.Context, token string) *model.AppError + RevokeAllSessions(c *request.Context, userID string) *model.AppError + RevokeSession(c *request.Context, session *model.Session) *model.AppError + RevokeSessionById(c *request.Context, sessionID string) *model.AppError + RevokeSessionsForDeviceId(c *request.Context, userID string, deviceID string, currentSessionId string) *model.AppError + RevokeUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError RolesGrantPermission(roleNames []string, permissionId string) bool Saml() einterfaces.SamlInterface SanitizePostListMetadataForUser(c request.CTX, postList *model.PostList, userID string) (*model.PostList, *model.AppError) @@ -1097,7 +1097,7 @@ type AppIface interface { SwitchEmailToLdap(c *request.Context, email, password, code, ldapLoginId, ldapPassword string) (string, *model.AppError) SwitchEmailToOAuth(c *request.Context, w http.ResponseWriter, r *http.Request, email, password, code, service string) (string, *model.AppError) SwitchLdapToEmail(c *request.Context, ldapPassword, code, email, newPassword string) (string, *model.AppError) - SwitchOAuthToEmail(email, password, requesterId string) (string, *model.AppError) + SwitchOAuthToEmail(c *request.Context, email, password, requesterId string) (string, *model.AppError) TeamMembersToRemove(teamID *string) ([]*model.TeamMember, *model.AppError) TelemetryId() string TestElasticsearch(cfg *model.Config) *model.AppError @@ -1111,7 +1111,7 @@ type AppIface interface { TotalWebsocketConnections() int TriggerWebhook(c request.CTX, payload *model.OutgoingWebhookPayload, hook *model.OutgoingWebhook, post *model.Post, channel *model.Channel) UnregisterPluginCommand(pluginID, teamID, trigger string) - UpdateActive(c request.CTX, user *model.User, active bool) (*model.User, *model.AppError) + UpdateActive(c *request.Context, user *model.User, active bool) (*model.User, *model.AppError) UpdateChannelMemberNotifyProps(c request.CTX, data map[string]string, channelID string, userID string) (*model.ChannelMember, *model.AppError) UpdateChannelMemberRoles(c request.CTX, channelID string, userID string, newRoles string) (*model.ChannelMember, *model.AppError) UpdateChannelMemberSchemeRoles(c request.CTX, channelID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.ChannelMember, *model.AppError) @@ -1146,8 +1146,8 @@ type AppIface interface { UpdateSidebarCategories(c request.CTX, userID, teamID string, categories []*model.SidebarCategoryWithChannels) ([]*model.SidebarCategoryWithChannels, *model.AppError) UpdateSidebarCategoryOrder(c request.CTX, userID, teamID string, categoryOrder []string) *model.AppError UpdateTeam(team *model.Team) (*model.Team, *model.AppError) - UpdateTeamMemberRoles(teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) - UpdateTeamMemberSchemeRoles(teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) + UpdateTeamMemberRoles(c request.CTX, teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) + UpdateTeamMemberSchemeRoles(c request.CTX, teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) UpdateTeamPrivacy(teamID string, teamType string, allowOpenInvite bool) *model.AppError UpdateTeamScheme(team *model.Team) (*model.Team, *model.AppError) UpdateThreadFollowForUser(userID, teamID, threadID string, state bool) *model.AppError @@ -1156,7 +1156,7 @@ type AppIface interface { UpdateThreadReadForUserByPost(c request.CTX, currentSessionId, userID, teamID, threadID, postID string) (*model.ThreadResponse, *model.AppError) UpdateThreadsReadForUser(userID, teamID string) *model.AppError UpdateUser(c request.CTX, user *model.User, sendNotifications bool) (*model.User, *model.AppError) - UpdateUserActive(c request.CTX, userID string, active bool) *model.AppError + UpdateUserActive(c *request.Context, userID string, active bool) *model.AppError UpdateUserAsUser(c request.CTX, user *model.User, asAdmin bool) (*model.User, *model.AppError) UpdateUserAuth(userID string, userAuth *model.UserAuth) (*model.UserAuth, *model.AppError) UpdateUserRoles(c request.CTX, userID string, newRoles string, sendWebSocketEvent bool) (*model.User, *model.AppError) @@ -1169,7 +1169,7 @@ type AppIface interface { UpsertGroupMembers(groupID string, userIDs []string) ([]*model.GroupMember, *model.AppError) UpsertGroupSyncable(groupSyncable *model.GroupSyncable) (*model.GroupSyncable, *model.AppError) UserAlreadyNotifiedOnRequiredFeature(user string, feature model.MattermostFeature) bool - UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) + UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) UserIsFirstAdmin(user *model.User) bool ValidateDesktopToken(token string, expiryTime int64) (*model.User, *model.AppError) VerifyEmailFromToken(c request.CTX, userSuppliedTokenString string) *model.AppError diff --git a/server/channels/app/authorization.go b/server/channels/app/authorization.go index 48543d2d2c..744f460bde 100644 --- a/server/channels/app/authorization.go +++ b/server/channels/app/authorization.go @@ -281,11 +281,11 @@ func (a *App) HasPermissionTo(askingUserId string, permission *model.Permission) return a.RolesGrantPermission(roles, permission.Id) } -func (a *App) HasPermissionToTeam(askingUserId string, teamID string, permission *model.Permission) bool { +func (a *App) HasPermissionToTeam(c request.CTX, askingUserId string, teamID string, permission *model.Permission) bool { if teamID == "" || askingUserId == "" { return false } - teamMember, _ := a.GetTeamMember(teamID, askingUserId) + teamMember, _ := a.GetTeamMember(c, teamID, askingUserId) if teamMember != nil && teamMember.DeleteAt == 0 { if a.RolesGrantPermission(teamMember.GetRoles(), permission.Id) { return true @@ -310,13 +310,13 @@ func (a *App) HasPermissionToChannel(c request.CTX, askingUserId string, channel var channel *model.Channel channel, err = a.GetChannel(c, channelID) if err == nil { - return a.HasPermissionToTeam(askingUserId, channel.TeamId, permission) + return a.HasPermissionToTeam(c, askingUserId, channel.TeamId, permission) } return a.HasPermissionTo(askingUserId, permission) } -func (a *App) HasPermissionToChannelByPost(askingUserId string, postID string, permission *model.Permission) bool { +func (a *App) HasPermissionToChannelByPost(c request.CTX, askingUserId string, postID string, permission *model.Permission) bool { if channelMember, err := a.Srv().Store().Channel().GetMemberForPost(postID, askingUserId); err == nil { if a.RolesGrantPermission(channelMember.GetRoles(), permission.Id) { return true @@ -324,7 +324,7 @@ func (a *App) HasPermissionToChannelByPost(askingUserId string, postID string, p } if channel, err := a.Srv().Store().Channel().GetForPost(postID); err == nil { - return a.HasPermissionToTeam(askingUserId, channel.TeamId, permission) + return a.HasPermissionToTeam(c, askingUserId, channel.TeamId, permission) } return a.HasPermissionTo(askingUserId, permission) @@ -406,5 +406,5 @@ func (a *App) HasPermissionToReadChannel(c request.CTX, userID string, channel * if !*a.Config().TeamSettings.ExperimentalViewArchivedChannels && channel.DeleteAt != 0 { return false } - return a.HasPermissionToChannel(c, userID, channel.Id, model.PermissionReadChannelContent) || (channel.Type == model.ChannelTypeOpen && a.HasPermissionToTeam(userID, channel.TeamId, model.PermissionReadPublicChannel)) + return a.HasPermissionToChannel(c, userID, channel.Id, model.PermissionReadChannelContent) || (channel.Type == model.ChannelTypeOpen && a.HasPermissionToTeam(c, userID, channel.TeamId, model.PermissionReadPublicChannel)) } diff --git a/server/channels/app/authorization_test.go b/server/channels/app/authorization_test.go index 6eb58e25e4..711988665f 100644 --- a/server/channels/app/authorization_test.go +++ b/server/channels/app/authorization_test.go @@ -57,17 +57,17 @@ func TestHasPermissionToTeam(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - assert.True(t, th.App.HasPermissionToTeam(th.BasicUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.True(t, th.App.HasPermissionToTeam(th.Context, th.BasicUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) th.RemoveUserFromTeam(th.BasicUser, th.BasicTeam) - assert.False(t, th.App.HasPermissionToTeam(th.BasicUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.False(t, th.App.HasPermissionToTeam(th.Context, th.BasicUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) - assert.True(t, th.App.HasPermissionToTeam(th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.True(t, th.App.HasPermissionToTeam(th.Context, th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) th.LinkUserToTeam(th.SystemAdminUser, th.BasicTeam) - assert.True(t, th.App.HasPermissionToTeam(th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.True(t, th.App.HasPermissionToTeam(th.Context, th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) th.RemovePermissionFromRole(model.PermissionListTeamChannels.Id, model.TeamUserRoleId) - assert.True(t, th.App.HasPermissionToTeam(th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.True(t, th.App.HasPermissionToTeam(th.Context, th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) th.RemoveUserFromTeam(th.SystemAdminUser, th.BasicTeam) - assert.True(t, th.App.HasPermissionToTeam(th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) + assert.True(t, th.App.HasPermissionToTeam(th.Context, th.SystemAdminUser.Id, th.BasicTeam.Id, model.PermissionListTeamChannels)) } func TestSessionHasPermissionToChannel(t *testing.T) { @@ -332,7 +332,7 @@ func TestSessionHasPermissionToManageUserOrBot(t *testing.T) { func TestHasPermissionToCategory(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - session, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}}) + session, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}}) require.Nil(t, err) categories, err := th.App.GetSidebarCategoriesForTeamForUser(th.Context, th.BasicUser.Id, th.BasicTeam.Id) @@ -418,7 +418,7 @@ func TestSessionHasPermissionToGroup(t *testing.T) { th.RemovePermissionFromRole(permission.Id, groupRole.Name) } - session, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}, Roles: systemRole.Name}) + session, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}, Roles: systemRole.Name}) require.Nil(t, err) result := th.App.SessionHasPermissionToGroup(*session, group.Id, permission) diff --git a/server/channels/app/bot.go b/server/channels/app/bot.go index 8fdd4c3267..36a86c54c2 100644 --- a/server/channels/app/bot.go +++ b/server/channels/app/bot.go @@ -388,7 +388,7 @@ func (a *App) GetBots(options *model.BotGetOptions) (model.BotList, *model.AppEr } // UpdateBotActive marks a bot as active or inactive, along with its corresponding user. -func (a *App) UpdateBotActive(c request.CTX, botUserId string, active bool) (*model.Bot, *model.AppError) { +func (a *App) UpdateBotActive(c *request.Context, botUserId string, active bool) (*model.Bot, *model.AppError) { user, nErr := a.Srv().Store().User().Get(context.Background(), botUserId) if nErr != nil { var nfErr *store.ErrNotFound @@ -495,7 +495,7 @@ func (a *App) UpdateBotOwner(botUserId, newOwnerId string) (*model.Bot, *model.A } // disableUserBots disables all bots owned by the given user. -func (a *App) disableUserBots(c request.CTX, userID string) *model.AppError { +func (a *App) disableUserBots(c *request.Context, userID string) *model.AppError { perPage := 20 for { options := &model.BotGetOptions{ diff --git a/server/channels/app/channel.go b/server/channels/app/channel.go index 381d92c4fb..f4af8bdfd7 100644 --- a/server/channels/app/channel.go +++ b/server/channels/app/channel.go @@ -1586,7 +1586,7 @@ func (a *App) addUserToChannel(c request.CTX, user *model.User, channel *model.C // AddUserToChannel adds a user to a given channel. func (a *App) AddUserToChannel(c request.CTX, user *model.User, channel *model.Channel, skipTeamMemberIntegrityCheck bool) (*model.ChannelMember, *model.AppError) { if !skipTeamMemberIntegrityCheck { - teamMember, nErr := a.Srv().Store().Team().GetMember(context.Background(), channel.TeamId, user.Id) + teamMember, nErr := a.Srv().Store().Team().GetMember(c, channel.TeamId, user.Id) if nErr != nil { var nfErr *store.ErrNotFound switch { @@ -2534,7 +2534,7 @@ func (a *App) removeUserFromChannel(c request.CTX, userIDToRemove string, remove return err } if len(currentMembers) == 0 { - teamMember, err := a.GetTeamMember(channel.TeamId, userIDToRemove) + teamMember, err := a.GetTeamMember(c, channel.TeamId, userIDToRemove) if err != nil { return model.NewAppError("removeUserFromChannel", "api.team.remove_user_from_team.missing.app_error", nil, "", http.StatusBadRequest).Wrap(err) } diff --git a/server/channels/app/channel_category.go b/server/channels/app/channel_category.go index 05ddf6888b..cd4a67ae12 100644 --- a/server/channels/app/channel_category.go +++ b/server/channels/app/channel_category.go @@ -14,8 +14,8 @@ import ( "github.com/mattermost/mattermost/server/v8/channels/store" ) -func (a *App) createInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, *model.AppError) { - categories, nErr := a.Srv().Store().Channel().CreateInitialSidebarCategories(userID, opts) +func (a *App) createInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, *model.AppError) { + categories, nErr := a.Srv().Store().Channel().CreateInitialSidebarCategories(c, userID, opts) if nErr != nil { return nil, model.NewAppError("createInitialSidebarCategories", "app.channel.create_initial_sidebar_categories.internal_error", nil, "", http.StatusInternalServerError).Wrap(nErr) } @@ -28,7 +28,7 @@ func (a *App) GetSidebarCategoriesForTeamForUser(c request.CTX, userID, teamID s categories, err := a.Srv().Store().Channel().GetSidebarCategoriesForTeamForUser(userID, teamID) if err == nil && len(categories.Categories) == 0 { // A user must always have categories, so migration must not have happened yet, and we should run it ourselves - categories, appErr = a.createInitialSidebarCategories(userID, &store.SidebarCategorySearchOpts{ + categories, appErr = a.createInitialSidebarCategories(c, userID, &store.SidebarCategorySearchOpts{ TeamID: teamID, ExcludeTeam: false, }) @@ -55,7 +55,7 @@ func (a *App) GetSidebarCategories(c request.CTX, userID string, opts *store.Sid categories, err := a.Srv().Store().Channel().GetSidebarCategories(userID, opts) if err == nil && len(categories.Categories) == 0 { // A user must always have categories, so migration must not have happened yet, and we should run it ourselves - categories, appErr = a.createInitialSidebarCategories(userID, opts) + categories, appErr = a.createInitialSidebarCategories(c, userID, opts) if appErr != nil { return nil, appErr } diff --git a/server/channels/app/channel_test.go b/server/channels/app/channel_test.go index 1c9c61332a..802aaedf2d 100644 --- a/server/channels/app/channel_test.go +++ b/server/channels/app/channel_test.go @@ -709,7 +709,7 @@ func TestLeaveLastChannel(t *testing.T) { t.Run("Guest leaves not last channel", func(t *testing.T) { err = th.App.LeaveChannel(th.Context, townSquare.Id, guest.Id) require.Nil(t, err) - _, err = th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + _, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) assert.Nil(t, err, "It should maintain the team membership") }) @@ -718,7 +718,7 @@ func TestLeaveLastChannel(t *testing.T) { assert.Nil(t, err, "It should allow to remove a guest user from the default channel") _, err = th.App.GetChannelMember(th.Context, th.BasicChannel.Id, guest.Id) assert.NotNil(t, err) - _, err = th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + _, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) assert.Nil(t, err, "It should remove the team membership") }) } diff --git a/server/channels/app/command.go b/server/channels/app/command.go index 4cff1135e6..6693d907a7 100644 --- a/server/channels/app/command.go +++ b/server/channels/app/command.go @@ -32,7 +32,7 @@ var atMentionRegexp = regexp.MustCompile(`\B@[[:alnum:]][[:alnum:]\.\-_:]*`) type CommandProvider interface { GetTrigger() string GetCommand(a *App, T i18n.TranslateFunc) *model.Command - DoCommand(a *App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse + DoCommand(a *App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse } var commandProviders = make(map[string]CommandProvider) @@ -179,7 +179,7 @@ func (a *App) ListAllCommands(teamID string, T i18n.TranslateFunc) ([]*model.Com } // @openTracingParams args -func (a *App) ExecuteCommand(c request.CTX, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) { +func (a *App) ExecuteCommand(c *request.Context, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) { trigger := "" message := "" index := strings.IndexFunc(args.Command, unicode.IsSpace) @@ -279,7 +279,7 @@ func (a *App) MentionsToTeamMembers(c request.CTX, message, teamID string) model continue } - _, err := a.GetTeamMember(teamID, userFromTrimmed.Id) + _, err := a.GetTeamMember(c, teamID, userFromTrimmed.Id) if err != nil { // The user is not in the team, so we should ignore it return @@ -292,7 +292,7 @@ func (a *App) MentionsToTeamMembers(c request.CTX, message, teamID string) model return } - _, err := a.GetTeamMember(teamID, user.Id) + _, err := a.GetTeamMember(c, teamID, user.Id) if err != nil { // The user is not in the team, so we should ignore it return @@ -355,7 +355,7 @@ func (a *App) MentionsToPublicChannels(c request.CTX, message, teamID string) mo // tryExecuteBuiltInCommand attempts to run a built in command based on the given arguments. If no such command can be // found, returns nil for all arguments. -func (a *App) tryExecuteBuiltInCommand(c request.CTX, args *model.CommandArgs, trigger string, message string) (*model.Command, *model.CommandResponse) { +func (a *App) tryExecuteBuiltInCommand(c *request.Context, args *model.CommandArgs, trigger string, message string) (*model.Command, *model.CommandResponse) { provider := GetCommandProvider(trigger) if provider == nil { return nil, nil diff --git a/server/channels/app/command_autocomplete_test.go b/server/channels/app/command_autocomplete_test.go index 3270de4ac3..eec111b856 100644 --- a/server/channels/app/command_autocomplete_test.go +++ b/server/channels/app/command_autocomplete_test.go @@ -658,7 +658,7 @@ func (p *testCommandProvider) GetCommand(a *App, T i18n.TranslateFunc) *model.Co } } -func (p *testCommandProvider) DoCommand(a *App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (p *testCommandProvider) DoCommand(a *App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return &model.CommandResponse{ Text: "I do nothing!", ResponseType: model.CommandResponseTypeEphemeral, diff --git a/server/channels/app/expirynotify_test.go b/server/channels/app/expirynotify_test.go index f9dd9dd4a7..3843917b5b 100644 --- a/server/channels/app/expirynotify_test.go +++ b/server/channels/app/expirynotify_test.go @@ -55,7 +55,7 @@ func TestNotifySessionsExpired(t *testing.T) { } for _, d := range data { - _, err := th.App.CreateSession(&model.Session{ + _, err := th.App.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, DeviceId: d.deviceID, ExpiresAt: d.expiresAt, diff --git a/server/channels/app/import_functions.go b/server/channels/app/import_functions.go index 9a0362b4fe..00cade7baa 100644 --- a/server/channels/app/import_functions.go +++ b/server/channels/app/import_functions.go @@ -804,7 +804,7 @@ func (a *App) importUserTeams(c request.CTX, user *model.User, data *[]imports.U isAdminByTeamId = map[string]bool{} ) - existingMemberships, nErr := a.Srv().Store().Team().GetTeamsForUser(context.Background(), user.Id, "", true) + existingMemberships, nErr := a.Srv().Store().Team().GetTeamsForUser(c, user.Id, "", true) if nErr != nil { return model.NewAppError("importUserTeams", "app.team.get_members.app_error", nil, "", http.StatusInternalServerError).Wrap(nErr) } @@ -916,12 +916,12 @@ func (a *App) importUserTeams(c request.CTX, user *model.User, data *[]imports.U for _, member := range append(newMembers, oldMembers...) { if member.ExplicitRoles != rolesByTeamId[member.TeamId] { - if _, err = a.UpdateTeamMemberRoles(member.TeamId, user.Id, rolesByTeamId[member.TeamId]); err != nil { + if _, err = a.UpdateTeamMemberRoles(c, member.TeamId, user.Id, rolesByTeamId[member.TeamId]); err != nil { return err } } - a.UpdateTeamMemberSchemeRoles(member.TeamId, user.Id, isGuestByTeamId[member.TeamId], isUserByTeamId[member.TeamId], isAdminByTeamId[member.TeamId]) + a.UpdateTeamMemberSchemeRoles(c, member.TeamId, user.Id, isGuestByTeamId[member.TeamId], isUserByTeamId[member.TeamId], isAdminByTeamId[member.TeamId]) } for _, team := range allTeams { diff --git a/server/channels/app/import_functions_test.go b/server/channels/app/import_functions_test.go index 905bbe23c3..42e7f94b2d 100644 --- a/server/channels/app/import_functions_test.go +++ b/server/channels/app/import_functions_test.go @@ -1077,7 +1077,7 @@ func TestImportImportUser(t *testing.T) { user, appErr = th.App.GetUserByUsername(username) require.Nil(t, appErr, "Failed to get user from database.") - teamMember, appErr := th.App.GetTeamMember(team.Id, user.Id) + teamMember, appErr := th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, appErr, "Failed to get team member from database.") require.Equal(t, "team_user", teamMember.Roles) @@ -1136,7 +1136,7 @@ func TestImportImportUser(t *testing.T) { assert.Nil(t, appErr) // Check both member properties. - teamMember, appErr = th.App.GetTeamMember(team.Id, user.Id) + teamMember, appErr = th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, appErr, "Failed to get team member from database.") require.Equal(t, "team_user team_admin", teamMember.Roles) @@ -1452,7 +1452,7 @@ func TestImportImportUser(t *testing.T) { user, appErr = th.App.GetUserByUsername(*userData.Username) require.Nil(t, appErr, "Failed to get user from database.") - teamMember, appErr = th.App.GetTeamMember(team.Id, user.Id) + teamMember, appErr = th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, appErr, "Failed to get the team member") assert.True(t, teamMember.SchemeAdmin) @@ -1494,7 +1494,7 @@ func TestImportImportUser(t *testing.T) { user, appErr = th.App.GetUserByUsername(*deletedUserData.Username) require.Nil(t, appErr, "Failed to get user from database.") - teamMember, appErr = th.App.GetTeamMember(team.Id, user.Id) + teamMember, appErr = th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, appErr, "Failed to get the team member") assert.False(t, teamMember.SchemeAdmin) @@ -1536,7 +1536,7 @@ func TestImportImportUser(t *testing.T) { user, appErr = th.App.GetUserByUsername(*deletedGuestData.Username) require.Nil(t, appErr, "Failed to get user from database.") - teamMember, appErr = th.App.GetTeamMember(team.Id, user.Id) + teamMember, appErr = th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, appErr, "Failed to get the team member") assert.False(t, teamMember.SchemeAdmin) @@ -1737,7 +1737,7 @@ func TestImportUserTeams(t *testing.T) { } else { require.Nil(t, err) } - teamMembers, nErr := th.App.Srv().Store().Team().GetTeamsForUser(context.Background(), user.Id, "", true) + teamMembers, nErr := th.App.Srv().Store().Team().GetTeamsForUser(th.Context, user.Id, "", true) require.NoError(t, nErr) require.Len(t, teamMembers, tc.expectedUserTeams) if tc.expectedUserTeams == 1 { @@ -1880,7 +1880,7 @@ func TestImportUserChannels(t *testing.T) { for _, tc := range tt { t.Run(tc.name, func(t *testing.T) { user := th.CreateUser() - _, _, err := th.App.ch.srv.teamService.JoinUserToTeam(th.BasicTeam, user) + _, _, err := th.App.ch.srv.teamService.JoinUserToTeam(th.Context, th.BasicTeam, user) require.NoError(t, err) // Two times import must end with the same results diff --git a/server/channels/app/ldap.go b/server/channels/app/ldap.go index e6dd417240..fb4512fa3e 100644 --- a/server/channels/app/ldap.go +++ b/server/channels/app/ldap.go @@ -102,7 +102,7 @@ func (a *App) SwitchEmailToLdap(c *request.Context, email, password, code, ldapL return "", err } - if err := a.RevokeAllSessions(user.Id); err != nil { + if err := a.RevokeAllSessions(c, user.Id); err != nil { return "", err } @@ -155,7 +155,7 @@ func (a *App) SwitchLdapToEmail(c *request.Context, ldapPassword, code, email, n return "", err } - if err := a.RevokeAllSessions(user.Id); err != nil { + if err := a.RevokeAllSessions(c, user.Id); err != nil { return "", err } diff --git a/server/channels/app/login.go b/server/channels/app/login.go index 0d962c8099..6f25db5870 100644 --- a/server/channels/app/login.go +++ b/server/channels/app/login.go @@ -179,7 +179,7 @@ func (a *App) DoLogin(c *request.Context, w http.ResponseWriter, r *http.Request a.ch.srv.platform.SetSessionExpireInHours(session, *a.Config().ServiceSettings.SessionLengthMobileInHours) // A special case where we logout of all other sessions with the same Id - if err := a.RevokeSessionsForDeviceId(user.Id, deviceID, ""); err != nil { + if err := a.RevokeSessionsForDeviceId(c, user.Id, deviceID, ""); err != nil { err.StatusCode = http.StatusInternalServerError return err } @@ -208,7 +208,7 @@ func (a *App) DoLogin(c *request.Context, w http.ResponseWriter, r *http.Request } var err *model.AppError - if session, err = a.CreateSession(session); err != nil { + if session, err = a.CreateSession(c, session); err != nil { err.StatusCode = http.StatusInternalServerError return err } diff --git a/server/channels/app/notification.go b/server/channels/app/notification.go index ea9b9085dd..f2c76543f2 100644 --- a/server/channels/app/notification.go +++ b/server/channels/app/notification.go @@ -887,7 +887,7 @@ func (a *App) sendNoUsersNotifiedByGroupInChannel(c request.CTX, sender *model.U // sendOutOfChannelMentions sends an ephemeral post to the sender of a post if any of the given potential mentions // are outside of the post's channel. Returns whether or not an ephemeral post was sent. func (a *App) sendOutOfChannelMentions(c request.CTX, sender *model.User, post *model.Post, channel *model.Channel, potentialMentions []string) (bool, error) { - outOfChannelUsers, outOfGroupsUsers, err := a.filterOutOfChannelMentions(sender, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupsUsers, err := a.filterOutOfChannelMentions(c, sender, post, channel, potentialMentions) if err != nil { return false, err } @@ -901,10 +901,10 @@ func (a *App) sendOutOfChannelMentions(c request.CTX, sender *model.User, post * return true, nil } -func (a *App) FilterUsersByVisible(viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) { +func (a *App) FilterUsersByVisible(c request.CTX, viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) { result := []*model.User{} for _, user := range otherUsers { - canSee, err := a.UserCanSeeOtherUser(viewer.Id, user.Id) + canSee, err := a.UserCanSeeOtherUser(c, viewer.Id, user.Id) if err != nil { return nil, err } @@ -915,7 +915,7 @@ func (a *App) FilterUsersByVisible(viewer *model.User, otherUsers []*model.User) return result, nil } -func (a *App) filterOutOfChannelMentions(sender *model.User, post *model.Post, channel *model.Channel, potentialMentions []string) ([]*model.User, []*model.User, error) { +func (a *App) filterOutOfChannelMentions(c request.CTX, sender *model.User, post *model.Post, channel *model.Channel, potentialMentions []string) ([]*model.User, []*model.User, error) { if post.IsSystemMessage() { return nil, nil, nil } @@ -936,7 +936,7 @@ func (a *App) filterOutOfChannelMentions(sender *model.User, post *model.Post, c // Filter out inactive users and bots allUsers := model.UserSlice(users).FilterByActive(true) allUsers = allUsers.FilterWithoutBots() - allUsers, appErr := a.FilterUsersByVisible(sender, allUsers) + allUsers, appErr := a.FilterUsersByVisible(c, sender, allUsers) if appErr != nil { return nil, nil, appErr } diff --git a/server/channels/app/notification_push_test.go b/server/channels/app/notification_push_test.go index dff3655099..8775c2f6ac 100644 --- a/server/channels/app/notification_push_test.go +++ b/server/channels/app/notification_push_test.go @@ -1067,7 +1067,7 @@ func TestBuildPushNotificationMessageMentions(t *testing.T) { func TestSendPushNotifications(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - _, err := th.App.CreateSession(&model.Session{ + _, err := th.App.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, DeviceId: "test", ExpiresAt: model.GetMillis() + 100000, @@ -1407,14 +1407,14 @@ func TestAllPushNotifications(t *testing.T) { var testData []userSession for i := 0; i < 10; i++ { u := th.CreateUser() - sess, err := th.App.CreateSession(&model.Session{ + sess, err := th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, }) require.Nil(t, err) // We don't need to track the 2nd session. - _, err = th.App.CreateSession(&model.Session{ + _, err = th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, diff --git a/server/channels/app/notification_test.go b/server/channels/app/notification_test.go index aeb60bbd19..e0ebd504cd 100644 --- a/server/channels/app/notification_test.go +++ b/server/channels/app/notification_test.go @@ -319,7 +319,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{user2.Username, user3.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, channel, potentialMentions) assert.NoError(t, err) assert.Len(t, outOfChannelUsers, 2) @@ -332,7 +332,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{user2.Username, user3.Username, user4.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(guest, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, guest, post, channel, potentialMentions) require.NoError(t, err) require.Len(t, outOfChannelUsers, 1) @@ -346,7 +346,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { } potentialMentions := []string{user2.Username, user3.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, channel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -360,7 +360,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { } potentialMentions := []string{user2.Username, user3.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, directChannel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, directChannel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -374,7 +374,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { } potentialMentions := []string{user2.Username, user3.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, groupChannel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, groupChannel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -389,7 +389,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{inactiveUser.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, channel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -403,7 +403,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{botUser.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, channel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -414,7 +414,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{"foo", "bar"} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, channel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, channel, potentialMentions) assert.NoError(t, err) assert.Nil(t, outOfChannelUsers) @@ -448,7 +448,7 @@ func TestFilterOutOfChannelMentions(t *testing.T) { post := &model.Post{} potentialMentions := []string{nonChannelMember.Username, nonGroupMember.Username} - outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(user1, post, constrainedChannel, potentialMentions) + outOfChannelUsers, outOfGroupUsers, err := th.App.filterOutOfChannelMentions(th.Context, user1, post, constrainedChannel, potentialMentions) assert.NoError(t, err) assert.Len(t, outOfChannelUsers, 1) diff --git a/server/channels/app/oauth.go b/server/channels/app/oauth.go index 86d81abe21..122dfe2ab6 100644 --- a/server/channels/app/oauth.go +++ b/server/channels/app/oauth.go @@ -146,8 +146,8 @@ func (a *App) GetOAuthAppsByCreator(userID string, page, perPage int) ([]*model. return oauthApps, nil } -func (a *App) GetOAuthImplicitRedirect(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { - session, err := a.GetOAuthAccessTokenForImplicitFlow(userID, authRequest) +func (a *App) GetOAuthImplicitRedirect(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { + session, err := a.GetOAuthAccessTokenForImplicitFlow(c, userID, authRequest) if err != nil { return "", err } @@ -184,7 +184,7 @@ func (a *App) GetOAuthCodeRedirect(userID string, authRequest *model.AuthorizeRe return uri.String(), nil } -func (a *App) AllowOAuthAppAccessToUser(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { +func (a *App) AllowOAuthAppAccessToUser(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { if !*a.Config().ServiceSettings.EnableOAuthServiceProvider { return "", model.NewAppError("AllowOAuthAppAccessToUser", "api.oauth.allow_oauth.turn_off.app_error", nil, "", http.StatusNotImplemented) } @@ -214,7 +214,7 @@ func (a *App) AllowOAuthAppAccessToUser(userID string, authRequest *model.Author case model.AuthCodeResponseType: redirectURI, err = a.GetOAuthCodeRedirect(userID, authRequest) case model.ImplicitResponseType: - redirectURI, err = a.GetOAuthImplicitRedirect(userID, authRequest) + redirectURI, err = a.GetOAuthImplicitRedirect(c, userID, authRequest) default: return authRequest.RedirectURI + "?error=unsupported_response_type&state=" + authRequest.State, nil } @@ -240,7 +240,7 @@ func (a *App) AllowOAuthAppAccessToUser(userID string, authRequest *model.Author return redirectURI, nil } -func (a *App) GetOAuthAccessTokenForImplicitFlow(userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) { +func (a *App) GetOAuthAccessTokenForImplicitFlow(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) { if !*a.Config().ServiceSettings.EnableOAuthServiceProvider { return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.disabled.app_error", nil, "", http.StatusNotImplemented) } @@ -255,7 +255,7 @@ func (a *App) GetOAuthAccessTokenForImplicitFlow(userID string, authRequest *mod return nil, err } - session, err := a.newSession(oauthApp, user) + session, err := a.newSession(c, oauthApp, user) if err != nil { return nil, err } @@ -269,7 +269,7 @@ func (a *App) GetOAuthAccessTokenForImplicitFlow(userID string, authRequest *mod return session, nil } -func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, code, secret, refreshToken string) (*model.AccessResponse, *model.AppError) { +func (a *App) GetOAuthAccessTokenForCodeFlow(c *request.Context, clientId, grantType, redirectURI, code, secret, refreshToken string) (*model.AccessResponse, *model.AppError) { if !*a.Config().ServiceSettings.EnableOAuthServiceProvider { return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.disabled.app_error", nil, "", http.StatusNotImplemented) } @@ -321,7 +321,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, c if accessData != nil { if accessData.IsExpired() { var access *model.AccessResponse - access, err := a.newSessionUpdateToken(oauthApp, accessData, user) + access, err := a.newSessionUpdateToken(c, oauthApp, accessData, user) if err != nil { return nil, err } @@ -338,7 +338,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, c } else { var session *model.Session // Create a new session and return new access token - session, err := a.newSession(oauthApp, user) + session, err := a.newSession(c, oauthApp, user) if err != nil { return nil, err } @@ -372,7 +372,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, c return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound) } - access, err := a.newSessionUpdateToken(oauthApp, accessData, user) + access, err := a.newSessionUpdateToken(c, oauthApp, accessData, user) if err != nil { return nil, err } @@ -382,7 +382,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, c return accessRsp, nil } -func (a *App) newSession(app *model.OAuthApp, user *model.User) (*model.Session, *model.AppError) { +func (a *App) newSession(c *request.Context, app *model.OAuthApp, user *model.User) (*model.Session, *model.AppError) { // Set new token an session session := &model.Session{UserId: user.Id, Roles: user.Roles, IsOAuth: true} session.GenerateCSRF() @@ -393,7 +393,7 @@ func (a *App) newSession(app *model.OAuthApp, user *model.User) (*model.Session, session.AddProp(model.SessionPropOs, "OAuth2") session.AddProp(model.SessionPropBrowser, "OAuth2") - session, err := a.Srv().Store().Session().Save(session) + session, err := a.Srv().Store().Session().Save(c, session) if err != nil { return nil, model.NewAppError("newSession", "api.oauth.get_access_token.internal_session.app_error", nil, "", http.StatusInternalServerError) } @@ -403,13 +403,13 @@ func (a *App) newSession(app *model.OAuthApp, user *model.User) (*model.Session, return session, nil } -func (a *App) newSessionUpdateToken(app *model.OAuthApp, accessData *model.AccessData, user *model.User) (*model.AccessResponse, *model.AppError) { +func (a *App) newSessionUpdateToken(c *request.Context, app *model.OAuthApp, accessData *model.AccessData, user *model.User) (*model.AccessResponse, *model.AppError) { // Remove the previous session if err := a.Srv().Store().Session().Remove(accessData.Token); err != nil { mlog.Warn("error removing access data token from session", mlog.Err(err)) } - session, err := a.newSession(app, user) + session, err := a.newSession(c, app, user) if err != nil { return nil, err } @@ -493,7 +493,7 @@ func (a *App) GetAuthorizedAppsForUser(userID string, page, perPage int) ([]*mod return apps, nil } -func (a *App) DeauthorizeOAuthAppForUser(userID, appID string) *model.AppError { +func (a *App) DeauthorizeOAuthAppForUser(c *request.Context, userID, appID string) *model.AppError { if !*a.Config().ServiceSettings.EnableOAuthServiceProvider { return model.NewAppError("DeauthorizeOAuthAppForUser", "api.oauth.allow_oauth.turn_off.app_error", nil, "", http.StatusNotImplemented) } @@ -505,7 +505,7 @@ func (a *App) DeauthorizeOAuthAppForUser(userID, appID string) *model.AppError { } for _, ad := range accessData { - if err := a.RevokeAccessToken(ad.Token); err != nil { + if err := a.RevokeAccessToken(c, ad.Token); err != nil { return err } @@ -548,8 +548,8 @@ func (a *App) RegenerateOAuthAppSecret(app *model.OAuthApp) (*model.OAuthApp, *m return app, nil } -func (a *App) RevokeAccessToken(token string) *model.AppError { - if err := a.ch.srv.platform.RevokeAccessToken(token); err != nil { +func (a *App) RevokeAccessToken(c *request.Context, token string) *model.AppError { + if err := a.ch.srv.platform.RevokeAccessToken(c, token); err != nil { switch { case errors.Is(err, platform.GetTokenError): return model.NewAppError("RevokeAccessToken", "api.oauth.revoke_access_token.get.app_error", nil, "", http.StatusBadRequest).Wrap(err) @@ -678,7 +678,7 @@ func (a *App) CompleteSwitchWithOAuth(c *request.Context, service string, userDa return nil, model.NewAppError("CompleteSwitchWithOAuth", MissingAccountError, nil, "", http.StatusInternalServerError).Wrap(nErr) } - if err := a.RevokeAllSessions(user.Id); err != nil { + if err := a.RevokeAllSessions(c, user.Id); err != nil { return nil, err } @@ -969,7 +969,7 @@ func (a *App) SwitchEmailToOAuth(c *request.Context, w http.ResponseWriter, r *h return authURL, nil } -func (a *App) SwitchOAuthToEmail(email, password, requesterId string) (string, *model.AppError) { +func (a *App) SwitchOAuthToEmail(c *request.Context, email, password, requesterId string) (string, *model.AppError) { if a.Srv().License() != nil && !*a.Config().ServiceSettings.ExperimentalEnableAuthenticationTransfer { return "", model.NewAppError("oauthToEmail", "api.user.oauth_to_email.not_available.app_error", nil, "", http.StatusForbidden) } @@ -991,11 +991,11 @@ func (a *App) SwitchOAuthToEmail(email, password, requesterId string) (string, * a.Srv().Go(func() { if err := a.Srv().EmailService.SendSignInChangeEmail(user.Email, T("api.templates.signin_change_email.body.method_email"), user.Locale, a.GetSiteURL()); err != nil { - mlog.Error("error sending signin change email", mlog.Err(err)) + c.Logger().Error("error sending signin change email", mlog.Err(err)) } }) - if err := a.RevokeAllSessions(requesterId); err != nil { + if err := a.RevokeAllSessions(c, requesterId); err != nil { return "", err } diff --git a/server/channels/app/oauth_test.go b/server/channels/app/oauth_test.go index 16a526be58..2c8cecee93 100644 --- a/server/channels/app/oauth_test.go +++ b/server/channels/app/oauth_test.go @@ -49,26 +49,26 @@ func TestGetOAuthAccessTokenForImplicitFlow(t *testing.T) { State: "123", } - session, err := th.App.GetOAuthAccessTokenForImplicitFlow(th.BasicUser.Id, authRequest) + session, err := th.App.GetOAuthAccessTokenForImplicitFlow(th.Context, th.BasicUser.Id, authRequest) assert.Nil(t, err) assert.NotNil(t, session) th.App.UpdateConfig(func(cfg *model.Config) { *cfg.ServiceSettings.EnableOAuthServiceProvider = false }) - session, err = th.App.GetOAuthAccessTokenForImplicitFlow(th.BasicUser.Id, authRequest) + session, err = th.App.GetOAuthAccessTokenForImplicitFlow(th.Context, th.BasicUser.Id, authRequest) assert.NotNil(t, err, "should fail - oauth2 disabled") assert.Nil(t, session) th.App.UpdateConfig(func(cfg *model.Config) { *cfg.ServiceSettings.EnableOAuthServiceProvider = true }) authRequest.ClientId = "junk" - session, err = th.App.GetOAuthAccessTokenForImplicitFlow(th.BasicUser.Id, authRequest) + session, err = th.App.GetOAuthAccessTokenForImplicitFlow(th.Context, th.BasicUser.Id, authRequest) assert.NotNil(t, err, "should fail - bad client id") assert.Nil(t, session) authRequest.ClientId = oapp.Id - session, err = th.App.GetOAuthAccessTokenForImplicitFlow("junk", authRequest) + session, err = th.App.GetOAuthAccessTokenForImplicitFlow(th.Context, "junk", authRequest) assert.NotNil(t, err, "should fail - bad user id") assert.Nil(t, session) } @@ -85,9 +85,9 @@ func TestOAuthRevokeAccessToken(t *testing.T) { th.App.SetSessionExpireInHours(session, 24) var err *model.AppError - session, err = th.App.CreateSession(session) + session, err = th.App.CreateSession(th.Context, session) require.Nil(t, err) - err = th.App.RevokeAccessToken(session.Token) + err = th.App.RevokeAccessToken(th.Context, session.Token) require.NotNil(t, err, "Should have failed does not have an access token") require.Equal(t, http.StatusBadRequest, err.StatusCode) } @@ -116,7 +116,7 @@ func TestOAuthDeleteApp(t *testing.T) { session.IsOAuth = true th.App.ch.srv.platform.SetSessionExpireInHours(session, 24) - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) accessData := &model.AccessData{} accessData.Token = session.Token @@ -619,7 +619,7 @@ func TestDeauthorizeOAuthApp(t *testing.T) { redirectUrl, err := th.App.GetOAuthCodeRedirect(th.BasicUser.Id, authRequest) assert.Nil(t, err) - dErr := th.App.DeauthorizeOAuthAppForUser(th.BasicUser.Id, oapp.Id) + dErr := th.App.DeauthorizeOAuthAppForUser(th.Context, th.BasicUser.Id, oapp.Id) assert.Nil(t, dErr) uri, uErr := url.Parse(redirectUrl) @@ -670,7 +670,7 @@ func TestDeactivatedUserOAuthApp(t *testing.T) { _, appErr := th.App.UpdateActive(th.Context, th.BasicUser, false) require.Nil(t, appErr) - resp, accErr := th.App.GetOAuthAccessTokenForCodeFlow(oapp.Id, model.AccessTokenGrantType, oapp.CallbackUrls[0], code, oapp.ClientSecret, "") + resp, accErr := th.App.GetOAuthAccessTokenForCodeFlow(th.Context, oapp.Id, model.AccessTokenGrantType, oapp.CallbackUrls[0], code, oapp.ClientSecret, "") assert.Nil(t, resp) require.NotNil(t, accErr, "Should not get access token") require.Equal(t, http.StatusBadRequest, accErr.StatusCode) diff --git a/server/channels/app/opentracing/opentracing_layer.go b/server/channels/app/opentracing/opentracing_layer.go index 52fcbf5886..66d2bb7090 100644 --- a/server/channels/app/opentracing/opentracing_layer.go +++ b/server/channels/app/opentracing/opentracing_layer.go @@ -659,7 +659,7 @@ func (a *OpenTracingAppLayer) AdjustTeamsFromProductLimits(teamLimits *model.Tea return resultVar0 } -func (a *OpenTracingAppLayer) AllowOAuthAppAccessToUser(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { +func (a *OpenTracingAppLayer) AllowOAuthAppAccessToUser(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.AllowOAuthAppAccessToUser") @@ -671,7 +671,7 @@ func (a *OpenTracingAppLayer) AllowOAuthAppAccessToUser(userID string, authReque }() defer span.Finish() - resultVar0, resultVar1 := a.app.AllowOAuthAppAccessToUser(userID, authRequest) + resultVar0, resultVar1 := a.app.AllowOAuthAppAccessToUser(c, userID, authRequest) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -2461,7 +2461,7 @@ func (a *OpenTracingAppLayer) CreateScheme(scheme *model.Scheme) (*model.Scheme, return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) CreateSession(session *model.Session) (*model.Session, *model.AppError) { +func (a *OpenTracingAppLayer) CreateSession(c *request.Context, session *model.Session) (*model.Session, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.CreateSession") @@ -2473,7 +2473,7 @@ func (a *OpenTracingAppLayer) CreateSession(session *model.Session) (*model.Sess }() defer span.Finish() - resultVar0, resultVar1 := a.app.CreateSession(session) + resultVar0, resultVar1 := a.app.CreateSession(c, session) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -2857,7 +2857,7 @@ func (a *OpenTracingAppLayer) DeactivateMfa(userID string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) DeauthorizeOAuthAppForUser(userID string, appID string) *model.AppError { +func (a *OpenTracingAppLayer) DeauthorizeOAuthAppForUser(c *request.Context, userID string, appID string) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.DeauthorizeOAuthAppForUser") @@ -2869,7 +2869,7 @@ func (a *OpenTracingAppLayer) DeauthorizeOAuthAppForUser(userID string, appID st }() defer span.Finish() - resultVar0 := a.app.DeauthorizeOAuthAppForUser(userID, appID) + resultVar0 := a.app.DeauthorizeOAuthAppForUser(c, userID, appID) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -3593,7 +3593,7 @@ func (a *OpenTracingAppLayer) DeleteToken(token *model.Token) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) DemoteUserToGuest(c request.CTX, user *model.User) *model.AppError { +func (a *OpenTracingAppLayer) DemoteUserToGuest(c *request.Context, user *model.User) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.DemoteUserToGuest") @@ -3659,7 +3659,7 @@ func (a *OpenTracingAppLayer) DisablePlugin(id string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) DisableUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *OpenTracingAppLayer) DisableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.DisableUserAccessToken") @@ -3671,7 +3671,7 @@ func (a *OpenTracingAppLayer) DisableUserAccessToken(token *model.UserAccessToke }() defer span.Finish() - resultVar0 := a.app.DisableUserAccessToken(token) + resultVar0 := a.app.DisableUserAccessToken(c, token) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -4042,7 +4042,7 @@ func (a *OpenTracingAppLayer) EnablePlugin(id string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) EnableUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *OpenTracingAppLayer) EnableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.EnableUserAccessToken") @@ -4054,7 +4054,7 @@ func (a *OpenTracingAppLayer) EnableUserAccessToken(token *model.UserAccessToken }() defer span.Finish() - resultVar0 := a.app.EnableUserAccessToken(token) + resultVar0 := a.app.EnableUserAccessToken(c, token) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -4103,7 +4103,7 @@ func (a *OpenTracingAppLayer) EnvironmentConfig(filter func(reflect.StructField) return resultVar0 } -func (a *OpenTracingAppLayer) ExecuteCommand(c request.CTX, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) { +func (a *OpenTracingAppLayer) ExecuteCommand(c *request.Context, args *model.CommandArgs) (*model.CommandResponse, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.ExecuteCommand") @@ -4508,7 +4508,7 @@ func (a *OpenTracingAppLayer) FilterNonGroupTeamMembers(userIDs []string, team * return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) FilterUsersByVisible(viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) { +func (a *OpenTracingAppLayer) FilterUsersByVisible(c request.CTX, viewer *model.User, otherUsers []*model.User) ([]*model.User, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.FilterUsersByVisible") @@ -4520,7 +4520,7 @@ func (a *OpenTracingAppLayer) FilterUsersByVisible(viewer *model.User, otherUser }() defer span.Finish() - resultVar0, resultVar1 := a.app.FilterUsersByVisible(viewer, otherUsers) + resultVar0, resultVar1 := a.app.FilterUsersByVisible(c, viewer, otherUsers) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -7482,7 +7482,7 @@ func (a *OpenTracingAppLayer) GetNumberOfChannelsOnTeam(c request.CTX, teamID st return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetOAuthAccessTokenForCodeFlow(clientId string, grantType string, redirectURI string, code string, secret string, refreshToken string) (*model.AccessResponse, *model.AppError) { +func (a *OpenTracingAppLayer) GetOAuthAccessTokenForCodeFlow(c *request.Context, clientId string, grantType string, redirectURI string, code string, secret string, refreshToken string) (*model.AccessResponse, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetOAuthAccessTokenForCodeFlow") @@ -7494,7 +7494,7 @@ func (a *OpenTracingAppLayer) GetOAuthAccessTokenForCodeFlow(clientId string, gr }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, code, secret, refreshToken) + resultVar0, resultVar1 := a.app.GetOAuthAccessTokenForCodeFlow(c, clientId, grantType, redirectURI, code, secret, refreshToken) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -7504,7 +7504,7 @@ func (a *OpenTracingAppLayer) GetOAuthAccessTokenForCodeFlow(clientId string, gr return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetOAuthAccessTokenForImplicitFlow(userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) { +func (a *OpenTracingAppLayer) GetOAuthAccessTokenForImplicitFlow(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (*model.Session, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetOAuthAccessTokenForImplicitFlow") @@ -7516,7 +7516,7 @@ func (a *OpenTracingAppLayer) GetOAuthAccessTokenForImplicitFlow(userID string, }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetOAuthAccessTokenForImplicitFlow(userID, authRequest) + resultVar0, resultVar1 := a.app.GetOAuthAccessTokenForImplicitFlow(c, userID, authRequest) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -7614,7 +7614,7 @@ func (a *OpenTracingAppLayer) GetOAuthCodeRedirect(userID string, authRequest *m return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetOAuthImplicitRedirect(userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { +func (a *OpenTracingAppLayer) GetOAuthImplicitRedirect(c *request.Context, userID string, authRequest *model.AuthorizeRequest) (string, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetOAuthImplicitRedirect") @@ -7626,7 +7626,7 @@ func (a *OpenTracingAppLayer) GetOAuthImplicitRedirect(userID string, authReques }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetOAuthImplicitRedirect(userID, authRequest) + resultVar0, resultVar1 := a.app.GetOAuthImplicitRedirect(c, userID, authRequest) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -9234,7 +9234,7 @@ func (a *OpenTracingAppLayer) GetSession(token string) (*model.Session, *model.A return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetSessionById(sessionID string) (*model.Session, *model.AppError) { +func (a *OpenTracingAppLayer) GetSessionById(c *request.Context, sessionID string) (*model.Session, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetSessionById") @@ -9246,7 +9246,7 @@ func (a *OpenTracingAppLayer) GetSessionById(sessionID string) (*model.Session, }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetSessionById(sessionID) + resultVar0, resultVar1 := a.app.GetSessionById(c, sessionID) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -9273,7 +9273,7 @@ func (a *OpenTracingAppLayer) GetSessionLengthInMillis(session *model.Session) i return resultVar0 } -func (a *OpenTracingAppLayer) GetSessions(userID string) ([]*model.Session, *model.AppError) { +func (a *OpenTracingAppLayer) GetSessions(c *request.Context, userID string) ([]*model.Session, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetSessions") @@ -9285,7 +9285,7 @@ func (a *OpenTracingAppLayer) GetSessions(userID string) ([]*model.Session, *mod }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetSessions(userID) + resultVar0, resultVar1 := a.app.GetSessions(c, userID) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -9808,7 +9808,7 @@ func (a *OpenTracingAppLayer) GetTeamIdFromQuery(query url.Values) (string, *mod return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetTeamMember(teamID string, userID string) (*model.TeamMember, *model.AppError) { +func (a *OpenTracingAppLayer) GetTeamMember(c request.CTX, teamID string, userID string) (*model.TeamMember, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetTeamMember") @@ -9820,7 +9820,7 @@ func (a *OpenTracingAppLayer) GetTeamMember(teamID string, userID string) (*mode }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetTeamMember(teamID, userID) + resultVar0, resultVar1 := a.app.GetTeamMember(c, teamID, userID) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -9874,7 +9874,7 @@ func (a *OpenTracingAppLayer) GetTeamMembersByIds(teamID string, userIDs []strin return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetTeamMembersForUser(userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) { +func (a *OpenTracingAppLayer) GetTeamMembersForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetTeamMembersForUser") @@ -9886,7 +9886,7 @@ func (a *OpenTracingAppLayer) GetTeamMembersForUser(userID string, excludeTeamID }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetTeamMembersForUser(userID, excludeTeamID, includeDeleted) + resultVar0, resultVar1 := a.app.GetTeamMembersForUser(c, userID, excludeTeamID, includeDeleted) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -11223,7 +11223,7 @@ func (a *OpenTracingAppLayer) GetVerifyEmailToken(token string) (*model.Token, * return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetViewUsersRestrictions(userID string) (*model.ViewUsersRestrictions, *model.AppError) { +func (a *OpenTracingAppLayer) GetViewUsersRestrictions(c request.CTX, userID string) (*model.ViewUsersRestrictions, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetViewUsersRestrictions") @@ -11235,7 +11235,7 @@ func (a *OpenTracingAppLayer) GetViewUsersRestrictions(userID string) (*model.Vi }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetViewUsersRestrictions(userID) + resultVar0, resultVar1 := a.app.GetViewUsersRestrictions(c, userID) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -11456,7 +11456,7 @@ func (a *OpenTracingAppLayer) HasPermissionToChannel(c request.CTX, askingUserId return resultVar0 } -func (a *OpenTracingAppLayer) HasPermissionToChannelByPost(askingUserId string, postID string, permission *model.Permission) bool { +func (a *OpenTracingAppLayer) HasPermissionToChannelByPost(c request.CTX, askingUserId string, postID string, permission *model.Permission) bool { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.HasPermissionToChannelByPost") @@ -11468,7 +11468,7 @@ func (a *OpenTracingAppLayer) HasPermissionToChannelByPost(askingUserId string, }() defer span.Finish() - resultVar0 := a.app.HasPermissionToChannelByPost(askingUserId, postID, permission) + resultVar0 := a.app.HasPermissionToChannelByPost(c, askingUserId, postID, permission) return resultVar0 } @@ -11490,7 +11490,7 @@ func (a *OpenTracingAppLayer) HasPermissionToReadChannel(c request.CTX, userID s return resultVar0 } -func (a *OpenTracingAppLayer) HasPermissionToTeam(askingUserId string, teamID string, permission *model.Permission) bool { +func (a *OpenTracingAppLayer) HasPermissionToTeam(c request.CTX, askingUserId string, teamID string, permission *model.Permission) bool { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.HasPermissionToTeam") @@ -11502,7 +11502,7 @@ func (a *OpenTracingAppLayer) HasPermissionToTeam(askingUserId string, teamID st }() defer span.Finish() - resultVar0 := a.app.HasPermissionToTeam(askingUserId, teamID, permission) + resultVar0 := a.app.HasPermissionToTeam(c, askingUserId, teamID, permission) return resultVar0 } @@ -14387,7 +14387,7 @@ func (a *OpenTracingAppLayer) RestoreTeam(teamID string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) RestrictUsersGetByPermissions(userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) { +func (a *OpenTracingAppLayer) RestrictUsersGetByPermissions(c request.CTX, userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RestrictUsersGetByPermissions") @@ -14399,7 +14399,7 @@ func (a *OpenTracingAppLayer) RestrictUsersGetByPermissions(userID string, optio }() defer span.Finish() - resultVar0, resultVar1 := a.app.RestrictUsersGetByPermissions(userID, options) + resultVar0, resultVar1 := a.app.RestrictUsersGetByPermissions(c, userID, options) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -14409,7 +14409,7 @@ func (a *OpenTracingAppLayer) RestrictUsersGetByPermissions(userID string, optio return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) RestrictUsersSearchByPermissions(userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) { +func (a *OpenTracingAppLayer) RestrictUsersSearchByPermissions(c request.CTX, userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RestrictUsersSearchByPermissions") @@ -14421,7 +14421,7 @@ func (a *OpenTracingAppLayer) RestrictUsersSearchByPermissions(userID string, op }() defer span.Finish() - resultVar0, resultVar1 := a.app.RestrictUsersSearchByPermissions(userID, options) + resultVar0, resultVar1 := a.app.RestrictUsersSearchByPermissions(c, userID, options) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -14446,7 +14446,7 @@ func (a *OpenTracingAppLayer) ReturnSessionToPool(session *model.Session) { a.app.ReturnSessionToPool(session) } -func (a *OpenTracingAppLayer) RevokeAccessToken(token string) *model.AppError { +func (a *OpenTracingAppLayer) RevokeAccessToken(c *request.Context, token string) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeAccessToken") @@ -14458,7 +14458,7 @@ func (a *OpenTracingAppLayer) RevokeAccessToken(token string) *model.AppError { }() defer span.Finish() - resultVar0 := a.app.RevokeAccessToken(token) + resultVar0 := a.app.RevokeAccessToken(c, token) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -14468,7 +14468,7 @@ func (a *OpenTracingAppLayer) RevokeAccessToken(token string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) RevokeAllSessions(userID string) *model.AppError { +func (a *OpenTracingAppLayer) RevokeAllSessions(c *request.Context, userID string) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeAllSessions") @@ -14480,7 +14480,7 @@ func (a *OpenTracingAppLayer) RevokeAllSessions(userID string) *model.AppError { }() defer span.Finish() - resultVar0 := a.app.RevokeAllSessions(userID) + resultVar0 := a.app.RevokeAllSessions(c, userID) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -14490,7 +14490,7 @@ func (a *OpenTracingAppLayer) RevokeAllSessions(userID string) *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) RevokeSession(session *model.Session) *model.AppError { +func (a *OpenTracingAppLayer) RevokeSession(c *request.Context, session *model.Session) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeSession") @@ -14502,7 +14502,7 @@ func (a *OpenTracingAppLayer) RevokeSession(session *model.Session) *model.AppEr }() defer span.Finish() - resultVar0 := a.app.RevokeSession(session) + resultVar0 := a.app.RevokeSession(c, session) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -14512,7 +14512,7 @@ func (a *OpenTracingAppLayer) RevokeSession(session *model.Session) *model.AppEr return resultVar0 } -func (a *OpenTracingAppLayer) RevokeSessionById(sessionID string) *model.AppError { +func (a *OpenTracingAppLayer) RevokeSessionById(c *request.Context, sessionID string) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeSessionById") @@ -14524,7 +14524,7 @@ func (a *OpenTracingAppLayer) RevokeSessionById(sessionID string) *model.AppErro }() defer span.Finish() - resultVar0 := a.app.RevokeSessionById(sessionID) + resultVar0 := a.app.RevokeSessionById(c, sessionID) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -14534,7 +14534,7 @@ func (a *OpenTracingAppLayer) RevokeSessionById(sessionID string) *model.AppErro return resultVar0 } -func (a *OpenTracingAppLayer) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) *model.AppError { +func (a *OpenTracingAppLayer) RevokeSessionsForDeviceId(c *request.Context, userID string, deviceID string, currentSessionId string) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeSessionsForDeviceId") @@ -14546,7 +14546,7 @@ func (a *OpenTracingAppLayer) RevokeSessionsForDeviceId(userID string, deviceID }() defer span.Finish() - resultVar0 := a.app.RevokeSessionsForDeviceId(userID, deviceID, currentSessionId) + resultVar0 := a.app.RevokeSessionsForDeviceId(c, userID, deviceID, currentSessionId) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -14578,7 +14578,7 @@ func (a *OpenTracingAppLayer) RevokeSessionsFromAllUsers() *model.AppError { return resultVar0 } -func (a *OpenTracingAppLayer) RevokeUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *OpenTracingAppLayer) RevokeUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeUserAccessToken") @@ -14590,7 +14590,7 @@ func (a *OpenTracingAppLayer) RevokeUserAccessToken(token *model.UserAccessToken }() defer span.Finish() - resultVar0 := a.app.RevokeUserAccessToken(token) + resultVar0 := a.app.RevokeUserAccessToken(c, token) if resultVar0 != nil { span.LogFields(spanlog.Error(resultVar0)) @@ -16714,7 +16714,7 @@ func (a *OpenTracingAppLayer) SwitchLdapToEmail(c *request.Context, ldapPassword return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) SwitchOAuthToEmail(email string, password string, requesterId string) (string, *model.AppError) { +func (a *OpenTracingAppLayer) SwitchOAuthToEmail(c *request.Context, email string, password string, requesterId string) (string, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.SwitchOAuthToEmail") @@ -16726,7 +16726,7 @@ func (a *OpenTracingAppLayer) SwitchOAuthToEmail(email string, password string, }() defer span.Finish() - resultVar0, resultVar1 := a.app.SwitchOAuthToEmail(email, password, requesterId) + resultVar0, resultVar1 := a.app.SwitchOAuthToEmail(c, email, password, requesterId) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -17094,7 +17094,7 @@ func (a *OpenTracingAppLayer) UnregisterPluginCommand(pluginID string, teamID st a.app.UnregisterPluginCommand(pluginID, teamID, trigger) } -func (a *OpenTracingAppLayer) UpdateActive(c request.CTX, user *model.User, active bool) (*model.User, *model.AppError) { +func (a *OpenTracingAppLayer) UpdateActive(c *request.Context, user *model.User, active bool) (*model.User, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UpdateActive") @@ -17116,7 +17116,7 @@ func (a *OpenTracingAppLayer) UpdateActive(c request.CTX, user *model.User, acti return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) UpdateBotActive(c request.CTX, botUserId string, active bool) (*model.Bot, *model.AppError) { +func (a *OpenTracingAppLayer) UpdateBotActive(c *request.Context, botUserId string, active bool) (*model.Bot, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UpdateBotActive") @@ -17970,7 +17970,7 @@ func (a *OpenTracingAppLayer) UpdateTeam(team *model.Team) (*model.Team, *model. return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) UpdateTeamMemberRoles(teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) { +func (a *OpenTracingAppLayer) UpdateTeamMemberRoles(c request.CTX, teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UpdateTeamMemberRoles") @@ -17982,7 +17982,7 @@ func (a *OpenTracingAppLayer) UpdateTeamMemberRoles(teamID string, userID string }() defer span.Finish() - resultVar0, resultVar1 := a.app.UpdateTeamMemberRoles(teamID, userID, newRoles) + resultVar0, resultVar1 := a.app.UpdateTeamMemberRoles(c, teamID, userID, newRoles) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -17992,7 +17992,7 @@ func (a *OpenTracingAppLayer) UpdateTeamMemberRoles(teamID string, userID string return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) UpdateTeamMemberSchemeRoles(teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) { +func (a *OpenTracingAppLayer) UpdateTeamMemberSchemeRoles(c request.CTX, teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UpdateTeamMemberSchemeRoles") @@ -18004,7 +18004,7 @@ func (a *OpenTracingAppLayer) UpdateTeamMemberSchemeRoles(teamID string, userID }() defer span.Finish() - resultVar0, resultVar1 := a.app.UpdateTeamMemberSchemeRoles(teamID, userID, isSchemeGuest, isSchemeUser, isSchemeAdmin) + resultVar0, resultVar1 := a.app.UpdateTeamMemberSchemeRoles(c, teamID, userID, isSchemeGuest, isSchemeUser, isSchemeAdmin) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -18190,7 +18190,7 @@ func (a *OpenTracingAppLayer) UpdateUser(c request.CTX, user *model.User, sendNo return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) UpdateUserActive(c request.CTX, userID string, active bool) *model.AppError { +func (a *OpenTracingAppLayer) UpdateUserActive(c *request.Context, userID string, active bool) *model.AppError { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UpdateUserActive") @@ -18567,7 +18567,7 @@ func (a *OpenTracingAppLayer) UserAlreadyNotifiedOnRequiredFeature(user string, return resultVar0 } -func (a *OpenTracingAppLayer) UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) { +func (a *OpenTracingAppLayer) UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UserCanSeeOtherUser") @@ -18579,7 +18579,7 @@ func (a *OpenTracingAppLayer) UserCanSeeOtherUser(userID string, otherUserId str }() defer span.Finish() - resultVar0, resultVar1 := a.app.UserCanSeeOtherUser(userID, otherUserId) + resultVar0, resultVar1 := a.app.UserCanSeeOtherUser(c, userID, otherUserId) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) diff --git a/server/channels/app/permissions.go b/server/channels/app/permissions.go index b85ca03c57..ebca53824c 100644 --- a/server/channels/app/permissions.go +++ b/server/channels/app/permissions.go @@ -33,12 +33,12 @@ func (s *permissionsServiceWrapper) HasPermissionTo(userID string, permission *m return s.app.HasPermissionTo(userID, permission) } -func (s *permissionsServiceWrapper) HasPermissionToTeam(userID string, teamID string, permission *model.Permission) bool { - return s.app.HasPermissionToTeam(userID, teamID, permission) +func (s *permissionsServiceWrapper) HasPermissionToTeam(c *request.Context, userID string, teamID string, permission *model.Permission) bool { + return s.app.HasPermissionToTeam(c, userID, teamID, permission) } -func (s *permissionsServiceWrapper) HasPermissionToChannel(askingUserID string, channelID string, permission *model.Permission) bool { - return s.app.HasPermissionToChannel(request.EmptyContext(s.app.Log()), askingUserID, channelID, permission) +func (s *permissionsServiceWrapper) HasPermissionToChannel(c *request.Context, askingUserID string, channelID string, permission *model.Permission) bool { + return s.app.HasPermissionToChannel(c, askingUserID, channelID, permission) } func (s *permissionsServiceWrapper) RolesGrantPermission(roleNames []string, permissionId string) bool { diff --git a/server/channels/app/platform/helper_test.go b/server/channels/app/platform/helper_test.go index 015d50834b..28a9b632b0 100644 --- a/server/channels/app/platform/helper_test.go +++ b/server/channels/app/platform/helper_test.go @@ -12,6 +12,7 @@ import ( "github.com/stretchr/testify/mock" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" "github.com/mattermost/mattermost/server/v8/channels/store/storetest/mocks" "github.com/mattermost/mattermost/server/v8/channels/testlib" @@ -20,6 +21,7 @@ import ( ) type TestHelper struct { + Context *request.Context Service *PlatformService Suite SuiteIFace @@ -52,7 +54,7 @@ func (ms *mockSuite) GetSession(token string) (*model.Session, *model.AppError) return &model.Session{}, nil } func (ms *mockSuite) RolesGrantPermission(roleNames []string, permissionId string) bool { return true } -func (ms *mockSuite) UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) { +func (ms *mockSuite) UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) { return true, nil } @@ -158,6 +160,7 @@ func setupTestHelper(dbStore store.Store, enterprise bool, includeCacheLayer boo } th := &TestHelper{ + Context: request.TestContext(tb), Service: ps, Suite: &mockSuite{}, } diff --git a/server/channels/app/platform/license.go b/server/channels/app/platform/license.go index d71348d9e9..e237b8d90a 100644 --- a/server/channels/app/platform/license.go +++ b/server/channels/app/platform/license.go @@ -5,7 +5,6 @@ package platform import ( "bytes" - "context" "encoding/json" "fmt" "net/http" @@ -17,6 +16,7 @@ import ( "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/public/shared/mlog" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/jobs" "github.com/mattermost/mattermost/server/v8/channels/store/sqlstore" "github.com/mattermost/mattermost/server/v8/channels/utils" @@ -49,6 +49,8 @@ func (ps *PlatformService) License() *model.License { } func (ps *PlatformService) LoadLicense() { + c := request.EmptyContext(ps.logger) + // ENV var overrides all other sources of license. licenseStr := os.Getenv(LicenseEnv) if licenseStr != "" { @@ -97,7 +99,7 @@ func (ps *PlatformService) LoadLicense() { } } - record, nErr := ps.Store.License().Get(sqlstore.WithMaster(context.Background()), licenseId) + record, nErr := ps.Store.License().Get(sqlstore.RequestContextWithMaster(c), licenseId) if nErr != nil { ps.logger.Error("License key from https://mattermost.com required to unlock enterprise features.", mlog.Err(nErr)) ps.SetLicense(nil) diff --git a/server/channels/app/platform/mocks/SuiteIFace.go b/server/channels/app/platform/mocks/SuiteIFace.go index 7a147902c2..6bcadc2285 100644 --- a/server/channels/app/platform/mocks/SuiteIFace.go +++ b/server/channels/app/platform/mocks/SuiteIFace.go @@ -7,6 +7,8 @@ package mocks import ( model "github.com/mattermost/mattermost/server/public/model" mock "github.com/stretchr/testify/mock" + + request "github.com/mattermost/mattermost/server/public/shared/request" ) // SuiteIFace is an autogenerated mock type for the SuiteIFace type @@ -56,23 +58,23 @@ func (_m *SuiteIFace) RolesGrantPermission(roleNames []string, permissionId stri return r0 } -// UserCanSeeOtherUser provides a mock function with given fields: userID, otherUserId -func (_m *SuiteIFace) UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) { - ret := _m.Called(userID, otherUserId) +// UserCanSeeOtherUser provides a mock function with given fields: c, userID, otherUserId +func (_m *SuiteIFace) UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) { + ret := _m.Called(c, userID, otherUserId) var r0 bool var r1 *model.AppError - if rf, ok := ret.Get(0).(func(string, string) (bool, *model.AppError)); ok { - return rf(userID, otherUserId) + if rf, ok := ret.Get(0).(func(request.CTX, string, string) (bool, *model.AppError)); ok { + return rf(c, userID, otherUserId) } - if rf, ok := ret.Get(0).(func(string, string) bool); ok { - r0 = rf(userID, otherUserId) + if rf, ok := ret.Get(0).(func(request.CTX, string, string) bool); ok { + r0 = rf(c, userID, otherUserId) } else { r0 = ret.Get(0).(bool) } - if rf, ok := ret.Get(1).(func(string, string) *model.AppError); ok { - r1 = rf(userID, otherUserId) + if rf, ok := ret.Get(1).(func(request.CTX, string, string) *model.AppError); ok { + r1 = rf(c, userID, otherUserId) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) diff --git a/server/channels/app/platform/session.go b/server/channels/app/platform/session.go index 82ac01d26d..31836f6e81 100644 --- a/server/channels/app/platform/session.go +++ b/server/channels/app/platform/session.go @@ -4,13 +4,12 @@ package platform import ( - "context" "fmt" "time" "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/public/shared/mlog" - "github.com/mattermost/mattermost/server/v8/channels/store/sqlstore" + "github.com/mattermost/mattermost/server/public/shared/request" ) func (ps *PlatformService) ReturnSessionToPool(session *model.Session) { @@ -20,10 +19,10 @@ func (ps *PlatformService) ReturnSessionToPool(session *model.Session) { } } -func (ps *PlatformService) CreateSession(session *model.Session) (*model.Session, error) { +func (ps *PlatformService) CreateSession(c *request.Context, session *model.Session) (*model.Session, error) { session.Token = "" - session, err := ps.Store.Session().Save(session) + session, err := ps.Store.Session().Save(c, session) if err != nil { return nil, err } @@ -33,12 +32,12 @@ func (ps *PlatformService) CreateSession(session *model.Session) (*model.Session return session, nil } -func (ps *PlatformService) GetSessionContext(ctx context.Context, token string) (*model.Session, error) { - return ps.Store.Session().Get(ctx, token) +func (ps *PlatformService) GetSessionContext(c *request.Context, token string) (*model.Session, error) { + return ps.Store.Session().Get(c, token) } -func (ps *PlatformService) GetSessions(userID string) ([]*model.Session, error) { - return ps.Store.Session().GetSessions(userID) +func (ps *PlatformService) GetSessions(c *request.Context, userID string) ([]*model.Session, error) { + return ps.Store.Session().GetSessions(c, userID) } func (ps *PlatformService) AddSessionToCache(session *model.Session) { @@ -97,7 +96,7 @@ func (ps *PlatformService) ClearAllUsersSessionCache() { } } -func (ps *PlatformService) GetSession(token string) (*model.Session, error) { +func (ps *PlatformService) GetSession(c *request.Context, token string) (*model.Session, error) { var session = ps.sessionPool.Get().(*model.Session) if err := ps.sessionCache.Get(token, session); err == nil { if m := ps.metricsIFace; m != nil { @@ -113,11 +112,11 @@ func (ps *PlatformService) GetSession(token string) (*model.Session, error) { return session, nil } - return ps.GetSessionContext(sqlstore.WithMaster(context.Background()), token) + return ps.GetSessionContext(c, token) } -func (ps *PlatformService) GetSessionByID(sessionID string) (*model.Session, error) { - return ps.Store.Session().Get(context.Background(), sessionID) +func (ps *PlatformService) GetSessionByID(c *request.Context, sessionID string) (*model.Session, error) { + return ps.Store.Session().Get(c, sessionID) } func (ps *PlatformService) RevokeSessionsFromAllUsers() error { @@ -135,16 +134,16 @@ func (ps *PlatformService) RevokeSessionsFromAllUsers() error { return nil } -func (ps *PlatformService) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) error { - sessions, err := ps.Store.Session().GetSessions(userID) +func (ps *PlatformService) RevokeSessionsForDeviceId(c *request.Context, userID string, deviceID string, currentSessionId string) error { + sessions, err := ps.Store.Session().GetSessions(c, userID) if err != nil { return err } for _, session := range sessions { if session.DeviceId == deviceID && session.Id != currentSessionId { - mlog.Debug("Revoking sessionId for userId. Re-login with the same device Id", mlog.String("session_id", session.Id), mlog.String("user_id", userID)) - if err := ps.RevokeSession(session); err != nil { - mlog.Warn("Could not revoke session for device", mlog.String("device_id", deviceID), mlog.Err(err)) + c.Logger().Debug("Revoking sessionId for userId. Re-login with the same device Id", mlog.String("session_id", session.Id), mlog.String("user_id", userID)) + if err := ps.RevokeSession(c, session); err != nil { + c.Logger().Warn("Could not revoke session for device", mlog.String("device_id", deviceID), mlog.Err(err)) } } } @@ -152,9 +151,9 @@ func (ps *PlatformService) RevokeSessionsForDeviceId(userID string, deviceID str return nil } -func (ps *PlatformService) RevokeSession(session *model.Session) error { +func (ps *PlatformService) RevokeSession(c *request.Context, session *model.Session) error { if session.IsOAuth { - if err := ps.RevokeAccessToken(session.Token); err != nil { + if err := ps.RevokeAccessToken(c, session.Token); err != nil { return err } } else { @@ -168,8 +167,8 @@ func (ps *PlatformService) RevokeSession(session *model.Session) error { return nil } -func (ps *PlatformService) RevokeAccessToken(token string) error { - session, _ := ps.GetSession(token) +func (ps *PlatformService) RevokeAccessToken(c *request.Context, token string) error { + session, _ := ps.GetSession(c, token) defer ps.ReturnSessionToPool(session) @@ -223,8 +222,8 @@ func (ps *PlatformService) ExtendSessionExpiry(session *model.Session, newExpiry return nil } -func (ps *PlatformService) UpdateSessionsIsGuest(userID string, isGuest bool) error { - sessions, err := ps.GetSessions(userID) +func (ps *PlatformService) UpdateSessionsIsGuest(c *request.Context, userID string, isGuest bool) error { + sessions, err := ps.GetSessions(c, userID) if err != nil { return err } @@ -241,14 +240,14 @@ func (ps *PlatformService) UpdateSessionsIsGuest(userID string, isGuest bool) er return nil } -func (ps *PlatformService) RevokeAllSessions(userID string) error { - sessions, err := ps.Store.Session().GetSessions(userID) +func (ps *PlatformService) RevokeAllSessions(c *request.Context, userID string) error { + sessions, err := ps.Store.Session().GetSessions(c, userID) if err != nil { return fmt.Errorf("%s: %w", err.Error(), GetSessionError) } for _, session := range sessions { if session.IsOAuth { - ps.RevokeAccessToken(session.Token) + ps.RevokeAccessToken(c, session.Token) } else { if err := ps.Store.Session().Remove(session.Id); err != nil { return fmt.Errorf("%s: %w", err.Error(), DeleteSessionError) diff --git a/server/channels/app/platform/session_test.go b/server/channels/app/platform/session_test.go index f52dd09c1e..111410a55c 100644 --- a/server/channels/app/platform/session_test.go +++ b/server/channels/app/platform/session_test.go @@ -105,7 +105,7 @@ func TestOAuthRevokeAccessToken(t *testing.T) { th := Setup(t) defer th.TearDown() - err := th.Service.RevokeAccessToken(model.NewRandomString(16)) + err := th.Service.RevokeAccessToken(th.Context, model.NewRandomString(16)) require.Error(t, err, "Should have failed due to an incorrect token") session := &model.Session{} @@ -115,8 +115,8 @@ func TestOAuthRevokeAccessToken(t *testing.T) { session.Roles = model.SystemUserRoleId th.Service.SetSessionExpireInHours(session, 24) - session, _ = th.Service.CreateSession(session) - err = th.Service.RevokeAccessToken(session.Token) + session, _ = th.Service.CreateSession(th.Context, session) + err = th.Service.RevokeAccessToken(th.Context, session.Token) require.Error(t, err, "Should have failed does not have an access token") accessData := &model.AccessData{} @@ -129,6 +129,6 @@ func TestOAuthRevokeAccessToken(t *testing.T) { _, nErr := th.Service.Store.OAuth().SaveAccessData(accessData) require.NoError(t, nErr) - err = th.Service.RevokeAccessToken(accessData.Token) + err = th.Service.RevokeAccessToken(th.Context, accessData.Token) require.NoError(t, err) } diff --git a/server/channels/app/platform/web_conn.go b/server/channels/app/platform/web_conn.go index 3e04d1c342..6ac3e12057 100644 --- a/server/channels/app/platform/web_conn.go +++ b/server/channels/app/platform/web_conn.go @@ -24,6 +24,7 @@ import ( "github.com/mattermost/mattermost/server/public/plugin" "github.com/mattermost/mattermost/server/public/shared/i18n" "github.com/mattermost/mattermost/server/public/shared/mlog" + "github.com/mattermost/mattermost/server/public/shared/request" ) const ( @@ -701,7 +702,11 @@ func (wc *WebConn) ShouldSendEventToGuest(msg *model.WebSocketEvent) bool { return true } - canSee, err := wc.Suite.UserCanSeeOtherUser(wc.UserId, userID) + // In the future, other methods in WebConn will use a request.Context. + // For now, it's fine to create it here. + c := request.EmptyContext(wc.Platform.logger) + + canSee, err := wc.Suite.UserCanSeeOtherUser(c, wc.UserId, userID) if err != nil { mlog.Error("webhub.shouldSendEvent.", mlog.Err(err)) return false diff --git a/server/channels/app/platform/web_hub.go b/server/channels/app/platform/web_hub.go index 8630cb5bb4..62110cef0d 100644 --- a/server/channels/app/platform/web_hub.go +++ b/server/channels/app/platform/web_hub.go @@ -13,6 +13,7 @@ import ( "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/public/shared/mlog" + "github.com/mattermost/mattermost/server/public/shared/request" ) const ( @@ -23,7 +24,7 @@ const ( type SuiteIFace interface { GetSession(token string) (*model.Session, *model.AppError) RolesGrantPermission(roleNames []string, permissionId string) bool - UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) + UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) } type webConnActivityMessage struct { diff --git a/server/channels/app/platform/web_hub_test.go b/server/channels/app/platform/web_hub_test.go index 79128ae89c..2b0cc0b840 100644 --- a/server/channels/app/platform/web_hub_test.go +++ b/server/channels/app/platform/web_hub_test.go @@ -64,7 +64,7 @@ func TestHubStopWithMultipleConnections(t *testing.T) { s := httptest.NewServer(dummyWebsocketHandler(t)) defer s.Close() - session, err := th.Service.CreateSession(&model.Session{ + session, err := th.Service.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, }) require.NoError(t, err) @@ -88,7 +88,7 @@ func TestHubStopRaceCondition(t *testing.T) { // So we just use this quick hack for the test. s := httptest.NewServer(dummyWebsocketHandler(t)) - session, err := th.Service.CreateSession(&model.Session{ + session, err := th.Service.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, }) require.NoError(t, err) @@ -153,8 +153,8 @@ func TestHubSessionRevokeRace(t *testing.T) { mockSessionStore := mocks.SessionStore{} mockSessionStore.On("UpdateLastActivityAt", "id1", mock.Anything).Return(nil) - mockSessionStore.On("Save", mock.AnythingOfType("*model.Session")).Return(sess1, nil) - mockSessionStore.On("Get", mock.Anything, "id1").Return(sess1, nil) + mockSessionStore.On("Save", mock.AnythingOfType("*request.Context"), mock.AnythingOfType("*model.Session")).Return(sess1, nil) + mockSessionStore.On("Get", mock.AnythingOfType("*request.Context"), mock.Anything, "id1").Return(sess1, nil) mockSessionStore.On("Remove", "id1").Return(nil) mockStatusStore := mocks.StatusStore{} @@ -179,7 +179,7 @@ func TestHubSessionRevokeRace(t *testing.T) { s := httptest.NewServer(dummyWebsocketHandler(t)) defer s.Close() - session, err := th.Service.CreateSession(&model.Session{ + session, err := th.Service.CreateSession(th.Context, &model.Session{ UserId: "testid", }) require.NoError(t, err) @@ -464,7 +464,7 @@ func TestHubIsRegistered(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - session, err := th.Service.CreateSession(&model.Session{ + session, err := th.Service.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser.Id, }) require.NoError(t, err) @@ -488,7 +488,7 @@ func TestHubIsRegistered(t *testing.T) { assert.True(t, th.Service.SessionIsRegistered(*wc2.session.Load())) assert.True(t, th.Service.SessionIsRegistered(*wc3.session.Load())) - session4, err := th.Service.CreateSession(&model.Session{ + session4, err := th.Service.CreateSession(th.Context, &model.Session{ UserId: th.BasicUser2.Id, }) require.NoError(t, err) diff --git a/server/channels/app/plugin_api.go b/server/channels/app/plugin_api.go index 7c6f38cd3e..a393047662 100644 --- a/server/channels/app/plugin_api.go +++ b/server/channels/app/plugin_api.go @@ -213,7 +213,7 @@ func (api *PluginAPI) GetTeamMembers(teamID string, page, perPage int) ([]*model } func (api *PluginAPI) GetTeamMember(teamID, userID string) (*model.TeamMember, *model.AppError) { - return api.app.GetTeamMember(teamID, userID) + return api.app.GetTeamMember(api.ctx, teamID, userID) } func (api *PluginAPI) GetTeamMembersForUser(userID string, page int, perPage int) ([]*model.TeamMember, *model.AppError) { @@ -221,7 +221,7 @@ func (api *PluginAPI) GetTeamMembersForUser(userID string, page int, perPage int } func (api *PluginAPI) UpdateTeamMemberRoles(teamID, userID, newRoles string) (*model.TeamMember, *model.AppError) { - return api.app.UpdateTeamMemberRoles(teamID, userID, newRoles) + return api.app.UpdateTeamMemberRoles(api.ctx, teamID, userID, newRoles) } func (api *PluginAPI) GetTeamStats(teamID string) (*model.TeamStats, *model.AppError) { @@ -283,15 +283,15 @@ func (api *PluginAPI) DeletePreferencesForUser(userID string, preferences []mode } func (api *PluginAPI) GetSession(sessionID string) (*model.Session, *model.AppError) { - return api.app.GetSessionById(sessionID) + return api.app.GetSessionById(api.ctx, sessionID) } func (api *PluginAPI) CreateSession(session *model.Session) (*model.Session, *model.AppError) { - return api.app.CreateSession(session) + return api.app.CreateSession(api.ctx, session) } func (api *PluginAPI) ExtendSessionExpiry(sessionID string, expiresAt int64) *model.AppError { - session, err := api.app.ch.srv.platform.GetSessionByID(sessionID) + session, err := api.app.ch.srv.platform.GetSessionByID(api.ctx, sessionID) if err != nil { return model.NewAppError("extendSessionExpiry", "app.session.get_sessions.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } @@ -304,7 +304,7 @@ func (api *PluginAPI) ExtendSessionExpiry(sessionID string, expiresAt int64) *mo } func (api *PluginAPI) RevokeSession(sessionID string) *model.AppError { - return api.app.RevokeSessionById(sessionID) + return api.app.RevokeSessionById(api.ctx, sessionID) } func (api *PluginAPI) CreateUserAccessToken(token *model.UserAccessToken) (*model.UserAccessToken, *model.AppError) { @@ -317,7 +317,7 @@ func (api *PluginAPI) RevokeUserAccessToken(tokenID string) *model.AppError { return err } - return api.app.RevokeUserAccessToken(accessToken) + return api.app.RevokeUserAccessToken(api.ctx, accessToken) } func (api *PluginAPI) UpdateUser(user *model.User) (*model.User, *model.AppError) { @@ -963,7 +963,7 @@ func (api *PluginAPI) HasPermissionTo(userID string, permission *model.Permissio } func (api *PluginAPI) HasPermissionToTeam(userID, teamID string, permission *model.Permission) bool { - return api.app.HasPermissionToTeam(userID, teamID, permission) + return api.app.HasPermissionToTeam(api.ctx, userID, teamID, permission) } func (api *PluginAPI) HasPermissionToChannel(userID, channelID string, permission *model.Permission) bool { diff --git a/server/channels/app/plugin_api_test.go b/server/channels/app/plugin_api_test.go index 54b8ab6437..0b0d305236 100644 --- a/server/channels/app/plugin_api_test.go +++ b/server/channels/app/plugin_api_test.go @@ -1834,7 +1834,7 @@ func (*MockSlashCommandProvider) GetCommand(a *App, T i18n.TranslateFunc) *model } } -func (mscp *MockSlashCommandProvider) DoCommand(a *App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (mscp *MockSlashCommandProvider) DoCommand(a *App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { mscp.Args = args mscp.Message = message return &model.CommandResponse{ @@ -2226,14 +2226,14 @@ func TestSendPushNotification(t *testing.T) { var userSessions []userSession for i := 0; i < 3; i++ { u := th.CreateUser() - sess, err := th.App.CreateSession(&model.Session{ + sess, err := th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, }) require.Nil(t, err) // We don't need to track the 2nd session. - _, err = th.App.CreateSession(&model.Session{ + _, err = th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, diff --git a/server/channels/app/plugin_hooks_test.go b/server/channels/app/plugin_hooks_test.go index 02e8bf2e01..54ae6e5418 100644 --- a/server/channels/app/plugin_hooks_test.go +++ b/server/channels/app/plugin_hooks_test.go @@ -1473,14 +1473,14 @@ func TestHookNotificationWillBePushed(t *testing.T) { var userSessions []userSession for i := 0; i < 3; i++ { u := th.CreateUser() - sess, err := th.App.CreateSession(&model.Session{ + sess, err := th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, }) require.Nil(t, err) // We don't need to track the 2nd session. - _, err = th.App.CreateSession(&model.Session{ + _, err = th.App.CreateSession(th.Context, &model.Session{ UserId: u.Id, DeviceId: "deviceID" + u.Id, ExpiresAt: model.GetMillis() + 100000, diff --git a/server/channels/app/post.go b/server/channels/app/post.go index 58d8fc3f30..0ec24baf96 100644 --- a/server/channels/app/post.go +++ b/server/channels/app/post.go @@ -2205,9 +2205,9 @@ func (a *App) GetPostInfo(c request.CTX, postID string) (*model.PostInfo, *model } if team.Type == model.TeamOpen { - hasPermissionToAccessTeam = a.HasPermissionToTeam(userID, team.Id, model.PermissionJoinPublicTeams) + hasPermissionToAccessTeam = a.HasPermissionToTeam(c, userID, team.Id, model.PermissionJoinPublicTeams) } else if team.Type == model.TeamInvite { - hasPermissionToAccessTeam = a.HasPermissionToTeam(userID, team.Id, model.PermissionJoinPrivateTeams) + hasPermissionToAccessTeam = a.HasPermissionToTeam(c, userID, team.Id, model.PermissionJoinPrivateTeams) } } else { // This happens in case of DMs and GMs. @@ -2240,7 +2240,7 @@ func (a *App) GetPostInfo(c request.CTX, postID string) (*model.PostInfo, *model HasJoinedChannel: channelMemberErr == nil, } if team != nil { - _, teamMemberErr := a.GetTeamMember(team.Id, userID) + _, teamMemberErr := a.GetTeamMember(c, team.Id, userID) info.TeamId = team.Id info.TeamType = team.Type diff --git a/server/channels/app/post_test.go b/server/channels/app/post_test.go index 3e0584e356..be440579fe 100644 --- a/server/channels/app/post_test.go +++ b/server/channels/app/post_test.go @@ -2864,11 +2864,11 @@ func TestGetPostIfAuthorized(t *testing.T) { require.Nil(t, err) require.NotNil(t, post) - session1, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}}) + session1, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser.Id, Props: model.StringMap{}}) require.Nil(t, err) require.NotNil(t, session1) - session2, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser2.Id, Props: model.StringMap{}}) + session2, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser2.Id, Props: model.StringMap{}}) require.Nil(t, err) require.NotNil(t, session2) diff --git a/server/channels/app/session.go b/server/channels/app/session.go index 6c57d1b7e7..481ee62b7e 100644 --- a/server/channels/app/session.go +++ b/server/channels/app/session.go @@ -4,7 +4,6 @@ package app import ( - "context" "errors" "math" "net/http" @@ -12,14 +11,15 @@ import ( "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/public/shared/mlog" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/app/platform" "github.com/mattermost/mattermost/server/v8/channels/app/users" "github.com/mattermost/mattermost/server/v8/channels/audit" "github.com/mattermost/mattermost/server/v8/channels/store" ) -func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppError) { - session, err := a.ch.srv.platform.CreateSession(session) +func (a *App) CreateSession(c *request.Context, session *model.Session) (*model.Session, *model.AppError) { + session, err := a.ch.srv.platform.CreateSession(c, session) if err != nil { var invErr *store.ErrInvalidInput switch { @@ -64,10 +64,14 @@ func (a *App) GetRemoteClusterSession(token string, remoteId string) (*model.Ses } func (a *App) GetSession(token string) (*model.Session, *model.AppError) { + // Create a context as GetSession is used in a lot of places where no context is current present. + // Once more of the codebase is migrated to use a context, GetSession should accept one. + c := request.EmptyContext(a.Log()) + var session *model.Session // We intentionally skip the error check here, we only want to check if the token is valid. // If we don't have the session we are going to create one with the token eventually. - if session, _ = a.ch.srv.platform.GetSession(token); session != nil { + if session, _ = a.ch.srv.platform.GetSession(c, token); session != nil { if session.Token != token { return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]any{"Token": token, "Error": ""}, "session token is different from the one in DB", http.StatusUnauthorized) } @@ -79,7 +83,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { var appErr *model.AppError if session == nil || session.Id == "" { - session, appErr = a.createSessionForUserAccessToken(token) + session, appErr = a.createSessionForUserAccessToken(c, token) if appErr != nil { detailedError := "" statusCode := http.StatusUnauthorized @@ -87,7 +91,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { detailedError = appErr.Error() statusCode = appErr.StatusCode } else { - mlog.Warn("Error while creating session for user access token", mlog.Err(appErr)) + c.Logger().Warn("Error while creating session for user access token", mlog.Err(appErr)) } return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]any{"Token": token, "Error": detailedError}, "", statusCode) } @@ -111,9 +115,9 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { // gets called from (*WebConn).isMemberOfTeam and revoking a session involves // clearing the webconn cache, which needs the hub again. a.Srv().Go(func() { - err := a.RevokeSessionById(session.Id) + err := a.RevokeSessionById(c, session.Id) if err != nil { - mlog.Warn("Error while revoking session", mlog.Err(err)) + c.Logger().Warn("Error while revoking session", mlog.Err(err)) } }) return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]any{"Token": token, "Error": ""}, "idle timeout", http.StatusUnauthorized) @@ -123,8 +127,8 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { return session, nil } -func (a *App) GetSessions(userID string) ([]*model.Session, *model.AppError) { - sessions, err := a.ch.srv.platform.GetSessions(userID) +func (a *App) GetSessions(c *request.Context, userID string) ([]*model.Session, *model.AppError) { + sessions, err := a.ch.srv.platform.GetSessions(c, userID) if err != nil { return nil, model.NewAppError("GetSessions", "app.session.get_sessions.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } @@ -132,8 +136,8 @@ func (a *App) GetSessions(userID string) ([]*model.Session, *model.AppError) { return sessions, nil } -func (a *App) RevokeAllSessions(userID string) *model.AppError { - if err := a.ch.srv.platform.RevokeAllSessions(userID); err != nil { +func (a *App) RevokeAllSessions(c *request.Context, userID string) *model.AppError { + if err := a.ch.srv.platform.RevokeAllSessions(c, userID); err != nil { switch { case errors.Is(err, platform.GetSessionError): return model.NewAppError("RevokeAllSessions", "app.session.get_sessions.app_error", nil, "", http.StatusInternalServerError).Wrap(err) @@ -186,16 +190,16 @@ func (a *App) ClearSessionCacheForAllUsersSkipClusterSend() { a.Srv().Platform().ClearSessionCacheForAllUsersSkipClusterSend() } -func (a *App) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) *model.AppError { - if err := a.ch.srv.platform.RevokeSessionsForDeviceId(userID, deviceID, currentSessionId); err != nil { +func (a *App) RevokeSessionsForDeviceId(c *request.Context, userID string, deviceID string, currentSessionId string) *model.AppError { + if err := a.ch.srv.platform.RevokeSessionsForDeviceId(c, userID, deviceID, currentSessionId); err != nil { return model.NewAppError("RevokeSessionsForDeviceId", "app.session.get_sessions.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } return nil } -func (a *App) GetSessionById(sessionID string) (*model.Session, *model.AppError) { - session, err := a.ch.srv.platform.GetSessionByID(sessionID) +func (a *App) GetSessionById(c *request.Context, sessionID string) (*model.Session, *model.AppError) { + session, err := a.ch.srv.platform.GetSessionByID(c, sessionID) if err != nil { return nil, model.NewAppError("GetSessionById", "app.session.get.app_error", nil, "", http.StatusBadRequest).Wrap(err) } @@ -203,16 +207,17 @@ func (a *App) GetSessionById(sessionID string) (*model.Session, *model.AppError) return session, nil } -func (a *App) RevokeSessionById(sessionID string) *model.AppError { - session, err := a.GetSessionById(sessionID) +func (a *App) RevokeSessionById(c *request.Context, sessionID string) *model.AppError { + session, err := a.GetSessionById(c, sessionID) if err != nil { return model.NewAppError("RevokeSessionById", "app.session.get.app_error", nil, "", http.StatusBadRequest).Wrap(err) } - return a.RevokeSession(session) + + return a.RevokeSession(c, session) } -func (a *App) RevokeSession(session *model.Session) *model.AppError { - if err := a.ch.srv.platform.RevokeSession(session); err != nil { +func (a *App) RevokeSession(c *request.Context, session *model.Session) *model.AppError { + if err := a.ch.srv.platform.RevokeSession(c, session); err != nil { switch { case errors.Is(err, platform.DeleteSessionError): return model.NewAppError("RevokeSession", "app.session.remove.app_error", nil, "", http.StatusInternalServerError).Wrap(err) @@ -346,7 +351,7 @@ func (a *App) CreateUserAccessToken(token *model.UserAccessToken) (*model.UserAc return token, nil } -func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Session, *model.AppError) { +func (a *App) createSessionForUserAccessToken(c *request.Context, tokenString string) (*model.Session, *model.AppError) { token, nErr := a.Srv().Store().UserAccessToken().GetByToken(tokenString) if nErr != nil { return nil, model.NewAppError("createSessionForUserAccessToken", "app.user_access_token.invalid_or_missing", nil, "", http.StatusUnauthorized).Wrap(nErr) @@ -356,7 +361,7 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio return nil, model.NewAppError("createSessionForUserAccessToken", "app.user_access_token.invalid_or_missing", nil, "inactive_token", http.StatusUnauthorized) } - user, nErr := a.Srv().Store().User().Get(context.Background(), token.UserId) + user, nErr := a.Srv().Store().User().Get(c.Context(), token.UserId) if nErr != nil { var nfErr *store.ErrNotFound switch { @@ -394,7 +399,7 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio } a.ch.srv.platform.SetSessionExpireInHours(session, model.SessionUserAccessTokenExpiryHours) - session, nErr = a.Srv().Store().Session().Save(session) + session, nErr = a.Srv().Store().Session().Save(c, session) if nErr != nil { var invErr *store.ErrInvalidInput switch { @@ -410,9 +415,9 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio return session, nil } -func (a *App) RevokeUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *App) RevokeUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { var session *model.Session - session, _ = a.ch.srv.platform.GetSessionContext(context.Background(), token.Token) + session, _ = a.ch.srv.platform.GetSessionContext(c, token.Token) if err := a.Srv().Store().UserAccessToken().Delete(token.Id); err != nil { return model.NewAppError("RevokeUserAccessToken", "app.user_access_token.delete.app_error", nil, "", http.StatusInternalServerError).Wrap(err) @@ -422,12 +427,12 @@ func (a *App) RevokeUserAccessToken(token *model.UserAccessToken) *model.AppErro return nil } - return a.RevokeSession(session) + return a.RevokeSession(c, session) } -func (a *App) DisableUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *App) DisableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { var session *model.Session - session, _ = a.ch.srv.platform.GetSessionContext(context.Background(), token.Token) + session, _ = a.ch.srv.platform.GetSessionContext(c, token.Token) if err := a.Srv().Store().UserAccessToken().UpdateTokenDisable(token.Id); err != nil { return model.NewAppError("DisableUserAccessToken", "app.user_access_token.update_token_disable.app_error", nil, "", http.StatusInternalServerError).Wrap(err) @@ -437,12 +442,12 @@ func (a *App) DisableUserAccessToken(token *model.UserAccessToken) *model.AppErr return nil } - return a.RevokeSession(session) + return a.RevokeSession(c, session) } -func (a *App) EnableUserAccessToken(token *model.UserAccessToken) *model.AppError { +func (a *App) EnableUserAccessToken(c *request.Context, token *model.UserAccessToken) *model.AppError { var session *model.Session - session, _ = a.ch.srv.platform.GetSessionContext(context.Background(), token.Token) + session, _ = a.ch.srv.platform.GetSessionContext(c, token.Token) err := a.Srv().Store().UserAccessToken().UpdateTokenEnable(token.Id) if err != nil { diff --git a/server/channels/app/session_test.go b/server/channels/app/session_test.go index 34c068c509..074e2f0f73 100644 --- a/server/channels/app/session_test.go +++ b/server/channels/app/session_test.go @@ -4,7 +4,6 @@ package app import ( - "context" "fmt" "os" "testing" @@ -23,7 +22,7 @@ func TestGetSessionIdleTimeoutInMinutes(t *testing.T) { UserId: model.NewId(), } - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) th.App.Srv().SetLicense(model.NewTestLicense("compliance")) th.App.UpdateConfig(func(cfg *model.Config) { *cfg.ServiceSettings.SessionIdleTimeoutInMinutes = 5 }) @@ -51,7 +50,7 @@ func TestGetSessionIdleTimeoutInMinutes(t *testing.T) { IsOAuth: true, } - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) time = session.LastActivityAt - (1000 * 60 * 6) nErr = th.App.Srv().Store().Session().UpdateLastActivityAt(session.Id, time) require.NoError(t, nErr) @@ -66,7 +65,7 @@ func TestGetSessionIdleTimeoutInMinutes(t *testing.T) { } session.AddProp(model.SessionPropType, model.SessionTypeUserAccessToken) - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) time = session.LastActivityAt - (1000 * 60 * 6) nErr = th.App.Srv().Store().Session().UpdateLastActivityAt(session.Id, time) require.NoError(t, nErr) @@ -84,7 +83,7 @@ func TestGetSessionIdleTimeoutInMinutes(t *testing.T) { UserId: model.NewId(), } - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) time = session.LastActivityAt - (1000 * 60 * 6) nErr = th.App.Srv().Store().Session().UpdateLastActivityAt(session.Id, time) require.NoError(t, nErr) @@ -103,7 +102,7 @@ func TestUpdateSessionOnPromoteDemote(t *testing.T) { t.Run("Promote Guest to User updates the session", func(t *testing.T) { guest := th.CreateGuest() - session, err := th.App.CreateSession(&model.Session{UserId: guest.Id, Props: model.StringMap{model.SessionPropIsGuest: "true"}}) + session, err := th.App.CreateSession(th.Context, &model.Session{UserId: guest.Id, Props: model.StringMap{model.SessionPropIsGuest: "true"}}) require.Nil(t, err) rsession, err := th.App.GetSession(session.Token) @@ -127,7 +126,7 @@ func TestUpdateSessionOnPromoteDemote(t *testing.T) { t.Run("Demote User to Guest updates the session", func(t *testing.T) { user := th.CreateUser() - session, err := th.App.CreateSession(&model.Session{UserId: user.Id, Props: model.StringMap{model.SessionPropIsGuest: "false"}}) + session, err := th.App.CreateSession(th.Context, &model.Session{UserId: user.Id, Props: model.StringMap{model.SessionPropIsGuest: "false"}}) require.Nil(t, err) rsession, err := th.App.GetSession(session.Token) @@ -164,7 +163,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { UserId: model.NewId(), DeviceId: model.NewId(), } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -178,7 +177,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { model.UserAuthServiceIsMobile: "true", }, } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -193,7 +192,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { model.UserAuthServiceIsSaml: "true", }, } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -207,7 +206,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { model.UserAuthServiceIsOAuth: "true", }, } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -220,7 +219,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { Props: map[string]string{ model.UserAuthServiceIsSaml: "true", }} - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -231,7 +230,7 @@ func TestApp_GetSessionLengthInMillis(t *testing.T) { session := &model.Session{ UserId: model.NewId(), } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) sessionLength := th.App.GetSessionLengthInMillis(session) @@ -254,7 +253,7 @@ func TestApp_ExtendExpiryIfNeeded(t *testing.T) { UserId: model.NewId(), ExpiresAt: expires, } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) ok := th.App.ExtendSessionExpiryIfNeeded(session) @@ -268,7 +267,7 @@ func TestApp_ExtendExpiryIfNeeded(t *testing.T) { session := &model.Session{ UserId: model.NewId(), } - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) require.Nil(t, err) expires := model.GetMillis() + th.App.GetSessionLengthInMillis(session) @@ -298,7 +297,7 @@ func TestApp_ExtendExpiryIfNeeded(t *testing.T) { t.Run(fmt.Sprintf("%s session beyond threshold should update ExpiresAt based on feature enabled", test.name), func(t *testing.T) { th.App.UpdateConfig(func(cfg *model.Config) { *cfg.ServiceSettings.ExtendSessionLengthWithActivity = test.enabled }) - session, err := th.App.CreateSession(test.session) + session, err := th.App.CreateSession(th.Context, test.session) require.Nil(t, err) expires := model.GetMillis() + th.App.GetSessionLengthInMillis(session) - hourMillis @@ -317,12 +316,12 @@ func TestApp_ExtendExpiryIfNeeded(t *testing.T) { require.False(t, session.IsExpired()) // check cache was updated - cachedSession, errGet := th.App.ch.srv.platform.GetSession(session.Token) + cachedSession, errGet := th.App.ch.srv.platform.GetSession(th.Context, session.Token) require.NoError(t, errGet) require.Equal(t, session.ExpiresAt, cachedSession.ExpiresAt) // check database was updated. - storedSession, nErr := th.App.Srv().Store().Session().Get(context.Background(), session.Token) + storedSession, nErr := th.App.Srv().Store().Session().Get(th.Context, session.Token) require.NoError(t, nErr) require.Equal(t, session.ExpiresAt, storedSession.ExpiresAt) }) diff --git a/server/channels/app/slashcommands/auto_users.go b/server/channels/app/slashcommands/auto_users.go index 32c0afc41b..512b38b04e 100644 --- a/server/channels/app/slashcommands/auto_users.go +++ b/server/channels/app/slashcommands/auto_users.go @@ -122,7 +122,7 @@ func (cfg *AutoUserCreator) createRandomUser(c request.CTX) (*model.User, error) } if cfg.JoinTime != 0 { - teamMember, appErr := cfg.app.GetTeamMember(cfg.team.Id, ruser.Id) + teamMember, appErr := cfg.app.GetTeamMember(c, cfg.team.Id, ruser.Id) if appErr != nil { return nil, appErr } diff --git a/server/channels/app/slashcommands/command_away.go b/server/channels/app/slashcommands/command_away.go index 28c23689ac..9a89889d2e 100644 --- a/server/channels/app/slashcommands/command_away.go +++ b/server/channels/app/slashcommands/command_away.go @@ -34,7 +34,7 @@ func (*AwayProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*AwayProvider) DoCommand(a *app.App, _ request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*AwayProvider) DoCommand(a *app.App, _ *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { a.SetStatusAwayIfNeeded(args.UserId, true) return &model.CommandResponse{ResponseType: model.CommandResponseTypeEphemeral, Text: args.T("api.command_away.success")} diff --git a/server/channels/app/slashcommands/command_channel_header.go b/server/channels/app/slashcommands/command_channel_header.go index 0403e14274..ea2739ab00 100644 --- a/server/channels/app/slashcommands/command_channel_header.go +++ b/server/channels/app/slashcommands/command_channel_header.go @@ -35,7 +35,7 @@ func (*HeaderProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (*HeaderProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*HeaderProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { channel, err := a.GetChannel(c, args.ChannelId) if err != nil { return &model.CommandResponse{ diff --git a/server/channels/app/slashcommands/command_channel_purpose.go b/server/channels/app/slashcommands/command_channel_purpose.go index 62fc4e8417..137706772c 100644 --- a/server/channels/app/slashcommands/command_channel_purpose.go +++ b/server/channels/app/slashcommands/command_channel_purpose.go @@ -35,7 +35,7 @@ func (*PurposeProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comm } } -func (*PurposeProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*PurposeProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { channel, err := a.GetChannel(c, args.ChannelId) if err != nil { return &model.CommandResponse{ diff --git a/server/channels/app/slashcommands/command_channel_rename.go b/server/channels/app/slashcommands/command_channel_rename.go index c387ee2265..7b854aa988 100644 --- a/server/channels/app/slashcommands/command_channel_rename.go +++ b/server/channels/app/slashcommands/command_channel_rename.go @@ -38,7 +38,7 @@ func (*RenameProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (*RenameProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*RenameProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { channel, err := a.GetChannel(c, args.ChannelId) if err != nil { return &model.CommandResponse{ diff --git a/server/channels/app/slashcommands/command_code.go b/server/channels/app/slashcommands/command_code.go index 6b7f9acecf..ce87a20583 100644 --- a/server/channels/app/slashcommands/command_code.go +++ b/server/channels/app/slashcommands/command_code.go @@ -37,7 +37,7 @@ func (*CodeProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*CodeProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*CodeProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if message == "" { return &model.CommandResponse{Text: args.T("api.command_code.message.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } diff --git a/server/channels/app/slashcommands/command_custom_status.go b/server/channels/app/slashcommands/command_custom_status.go index 809e707e25..4c06e918a2 100644 --- a/server/channels/app/slashcommands/command_custom_status.go +++ b/server/channels/app/slashcommands/command_custom_status.go @@ -41,7 +41,7 @@ func (*CustomStatusProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model } } -func (*CustomStatusProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*CustomStatusProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if !*a.Config().TeamSettings.EnableCustomUserStatuses { return nil } diff --git a/server/channels/app/slashcommands/command_dnd.go b/server/channels/app/slashcommands/command_dnd.go index 1f613d910f..d80bab4270 100644 --- a/server/channels/app/slashcommands/command_dnd.go +++ b/server/channels/app/slashcommands/command_dnd.go @@ -34,7 +34,7 @@ func (*DndProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*DndProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*DndProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { a.SetStatusDoNotDisturb(args.UserId) return &model.CommandResponse{ResponseType: model.CommandResponseTypeEphemeral, Text: args.T("api.command_dnd.success")} diff --git a/server/channels/app/slashcommands/command_echo.go b/server/channels/app/slashcommands/command_echo.go index 493d18f430..e29c614ab3 100644 --- a/server/channels/app/slashcommands/command_echo.go +++ b/server/channels/app/slashcommands/command_echo.go @@ -42,7 +42,7 @@ func (*EchoProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*EchoProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*EchoProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if message == "" { return &model.CommandResponse{Text: args.T("api.command_echo.message.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } diff --git a/server/channels/app/slashcommands/command_expand_collapse.go b/server/channels/app/slashcommands/command_expand_collapse.go index 4ed28fde0c..2a15c79660 100644 --- a/server/channels/app/slashcommands/command_expand_collapse.go +++ b/server/channels/app/slashcommands/command_expand_collapse.go @@ -55,11 +55,11 @@ func (*CollapseProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Com } } -func (*ExpandProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*ExpandProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return setCollapsePreference(a, args, false) } -func (*CollapseProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*CollapseProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return setCollapsePreference(a, args, true) } diff --git a/server/channels/app/slashcommands/command_exportlink.go b/server/channels/app/slashcommands/command_exportlink.go index 81a75a58ca..6a98bdf38d 100644 --- a/server/channels/app/slashcommands/command_exportlink.go +++ b/server/channels/app/slashcommands/command_exportlink.go @@ -58,7 +58,7 @@ func (*ExportLinkProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.C } } -func (*ExportLinkProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*ExportLinkProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if !a.SessionHasPermissionTo(*c.Session(), model.PermissionManageSystem) { return &model.CommandResponse{ResponseType: model.CommandResponseTypeEphemeral, Text: args.T("api.command_exportlink.permission.app_error")} } diff --git a/server/channels/app/slashcommands/command_groupmsg.go b/server/channels/app/slashcommands/command_groupmsg.go index 43eb186876..a7bdde490f 100644 --- a/server/channels/app/slashcommands/command_groupmsg.go +++ b/server/channels/app/slashcommands/command_groupmsg.go @@ -39,7 +39,7 @@ func (*groupmsgProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Com } } -func (*groupmsgProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*groupmsgProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { targetUsers := map[string]*model.User{} targetUsersSlice := []string{args.UserId} invalidUsernames := []string{} @@ -55,7 +55,7 @@ func (*groupmsgProvider) DoCommand(a *app.App, c request.CTX, args *model.Comman continue } - canSee, err := a.UserCanSeeOtherUser(args.UserId, targetUser.Id) + canSee, err := a.UserCanSeeOtherUser(c, args.UserId, targetUser.Id) if err != nil { return &model.CommandResponse{Text: args.T("api.command_groupmsg.fail.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } diff --git a/server/channels/app/slashcommands/command_help.go b/server/channels/app/slashcommands/command_help.go index 87296867db..d1a13202eb 100644 --- a/server/channels/app/slashcommands/command_help.go +++ b/server/channels/app/slashcommands/command_help.go @@ -34,7 +34,7 @@ func (h *HelpProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (h *HelpProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (h *HelpProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { helpLink := *a.Config().SupportSettings.HelpLink if helpLink == "" { diff --git a/server/channels/app/slashcommands/command_invite.go b/server/channels/app/slashcommands/command_invite.go index e3ea1d82fa..aaf6f26b15 100644 --- a/server/channels/app/slashcommands/command_invite.go +++ b/server/channels/app/slashcommands/command_invite.go @@ -48,14 +48,14 @@ func (*InviteProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (i *InviteProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (i *InviteProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return &model.CommandResponse{ Text: i.doCommand(a, c, args, message), ResponseType: model.CommandResponseTypeEphemeral, } } -func (i *InviteProvider) doCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) string { +func (i *InviteProvider) doCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) string { if message == "" { return args.T("api.command_invite.missing_message.app_error") } diff --git a/server/channels/app/slashcommands/command_invite_people.go b/server/channels/app/slashcommands/command_invite_people.go index 702e13a48b..52b7e1ec74 100644 --- a/server/channels/app/slashcommands/command_invite_people.go +++ b/server/channels/app/slashcommands/command_invite_people.go @@ -42,12 +42,12 @@ func (*InvitePeopleProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model } } -func (*InvitePeopleProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { - if !a.HasPermissionToTeam(args.UserId, args.TeamId, model.PermissionInviteUser) { +func (*InvitePeopleProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { + if !a.HasPermissionToTeam(c, args.UserId, args.TeamId, model.PermissionInviteUser) { return &model.CommandResponse{Text: args.T("api.command_invite_people.permission.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } - if !a.HasPermissionToTeam(args.UserId, args.TeamId, model.PermissionAddUserToTeam) { + if !a.HasPermissionToTeam(c, args.UserId, args.TeamId, model.PermissionAddUserToTeam) { return &model.CommandResponse{Text: args.T("api.command_invite_people.permission.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } diff --git a/server/channels/app/slashcommands/command_join.go b/server/channels/app/slashcommands/command_join.go index 2bc6be9066..d6828f5320 100644 --- a/server/channels/app/slashcommands/command_join.go +++ b/server/channels/app/slashcommands/command_join.go @@ -37,7 +37,7 @@ func (*JoinProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*JoinProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*JoinProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { channelName := strings.ToLower(message) if strings.HasPrefix(message, "~") { diff --git a/server/channels/app/slashcommands/command_leave.go b/server/channels/app/slashcommands/command_leave.go index 81a1898e50..12215708a5 100644 --- a/server/channels/app/slashcommands/command_leave.go +++ b/server/channels/app/slashcommands/command_leave.go @@ -34,7 +34,7 @@ func (*LeaveProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comman } } -func (*LeaveProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*LeaveProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { var channel *model.Channel var noChannelErr *model.AppError if channel, noChannelErr = a.GetChannel(c, args.ChannelId); noChannelErr != nil { @@ -54,7 +54,7 @@ func (*LeaveProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandAr return &model.CommandResponse{Text: args.T("api.command_leave.fail.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } - member, err := a.GetTeamMember(team.Id, args.UserId) + member, err := a.GetTeamMember(c, team.Id, args.UserId) if err != nil || member.DeleteAt != 0 { return &model.CommandResponse{GotoLocation: args.SiteURL + "/"} } diff --git a/server/channels/app/slashcommands/command_loadtest.go b/server/channels/app/slashcommands/command_loadtest.go index 33f75daf59..5badd489d2 100644 --- a/server/channels/app/slashcommands/command_loadtest.go +++ b/server/channels/app/slashcommands/command_loadtest.go @@ -139,7 +139,7 @@ func (*LoadTestProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Com } } -func (lt *LoadTestProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (lt *LoadTestProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { commandResponse, err := lt.doCommand(a, c, args, message) if err != nil { c.Logger().Error("failed command /"+CmdTest, mlog.Err(err)) @@ -148,7 +148,7 @@ func (lt *LoadTestProvider) DoCommand(a *app.App, c request.CTX, args *model.Com return commandResponse } -func (lt *LoadTestProvider) doCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) (*model.CommandResponse, error) { +func (lt *LoadTestProvider) doCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) (*model.CommandResponse, error) { //This command is only available when EnableTesting is true if !*a.Config().ServiceSettings.EnableTesting { return &model.CommandResponse{}, nil @@ -291,7 +291,7 @@ func (*LoadTestProvider) SetupCommand(a *app.App, c request.CTX, args *model.Com return &model.CommandResponse{Text: "Created environment", ResponseType: model.CommandResponseTypeEphemeral}, nil } -func (*LoadTestProvider) ActivateUserCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) (*model.CommandResponse, error) { +func (*LoadTestProvider) ActivateUserCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) (*model.CommandResponse, error) { user_id := strings.TrimSpace(strings.TrimPrefix(message, "activate_user")) if err := a.UpdateUserActive(c, user_id, true); err != nil { return &model.CommandResponse{Text: "Failed to activate user", ResponseType: model.CommandResponseTypeEphemeral}, err @@ -300,7 +300,7 @@ func (*LoadTestProvider) ActivateUserCommand(a *app.App, c request.CTX, args *mo return &model.CommandResponse{Text: "Activated user", ResponseType: model.CommandResponseTypeEphemeral}, nil } -func (*LoadTestProvider) DeActivateUserCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) (*model.CommandResponse, error) { +func (*LoadTestProvider) DeActivateUserCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) (*model.CommandResponse, error) { user_id := strings.TrimSpace(strings.TrimPrefix(message, "deactivate_user")) if err := a.UpdateUserActive(c, user_id, false); err != nil { return &model.CommandResponse{Text: "Failed to deactivate user", ResponseType: model.CommandResponseTypeEphemeral}, err diff --git a/server/channels/app/slashcommands/command_logout.go b/server/channels/app/slashcommands/command_logout.go index 0263b83bb2..e30093f61e 100644 --- a/server/channels/app/slashcommands/command_logout.go +++ b/server/channels/app/slashcommands/command_logout.go @@ -35,7 +35,7 @@ func (*LogoutProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (*LogoutProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*LogoutProvider) DoCommand(a *app.App, _ *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { // Actual logout is handled client side. return &model.CommandResponse{GotoLocation: "/login"} } diff --git a/server/channels/app/slashcommands/command_marketplace.go b/server/channels/app/slashcommands/command_marketplace.go index 7e36c91ca8..28b73357aa 100644 --- a/server/channels/app/slashcommands/command_marketplace.go +++ b/server/channels/app/slashcommands/command_marketplace.go @@ -40,7 +40,7 @@ func (h *MarketplaceProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *mode } } -func (h *MarketplaceProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (h *MarketplaceProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { // This command is handled client-side and shouldn't hit the server. return &model.CommandResponse{ Text: args.T("api.command_marketplace.unsupported.app_error"), diff --git a/server/channels/app/slashcommands/command_me.go b/server/channels/app/slashcommands/command_me.go index 794f3bddc9..99e977687f 100644 --- a/server/channels/app/slashcommands/command_me.go +++ b/server/channels/app/slashcommands/command_me.go @@ -35,7 +35,7 @@ func (*MeProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command { } } -func (*MeProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*MeProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return &model.CommandResponse{ ResponseType: model.CommandResponseTypeInChannel, Type: model.PostTypeMe, diff --git a/server/channels/app/slashcommands/command_msg.go b/server/channels/app/slashcommands/command_msg.go index 297139ea12..d6edde9489 100644 --- a/server/channels/app/slashcommands/command_msg.go +++ b/server/channels/app/slashcommands/command_msg.go @@ -40,7 +40,7 @@ func (*msgProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*msgProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*msgProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { splitMessage := strings.SplitN(message, " ", 2) parsedMessage := "" @@ -62,7 +62,7 @@ func (*msgProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs return &model.CommandResponse{Text: args.T("api.command_msg.missing.app_error"), ResponseType: model.CommandResponseTypeEphemeral} } - canSee, err := a.UserCanSeeOtherUser(args.UserId, userProfile.Id) + canSee, err := a.UserCanSeeOtherUser(c, args.UserId, userProfile.Id) if err != nil { mlog.Error(err.Error()) return &model.CommandResponse{Text: args.T("api.command_msg.fail.app_error"), ResponseType: model.CommandResponseTypeEphemeral} diff --git a/server/channels/app/slashcommands/command_mute.go b/server/channels/app/slashcommands/command_mute.go index fd5e6894b1..51cd94fa35 100644 --- a/server/channels/app/slashcommands/command_mute.go +++ b/server/channels/app/slashcommands/command_mute.go @@ -37,7 +37,7 @@ func (*MuteProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*MuteProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*MuteProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { var channel *model.Channel var noChannelErr *model.AppError diff --git a/server/channels/app/slashcommands/command_offline.go b/server/channels/app/slashcommands/command_offline.go index 8b013084cd..9d1dfcfbe4 100644 --- a/server/channels/app/slashcommands/command_offline.go +++ b/server/channels/app/slashcommands/command_offline.go @@ -34,7 +34,7 @@ func (*OfflineProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comm } } -func (*OfflineProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*OfflineProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { a.SetStatusOffline(args.UserId, true) return &model.CommandResponse{ResponseType: model.CommandResponseTypeEphemeral, Text: args.T("api.command_offline.success")} diff --git a/server/channels/app/slashcommands/command_online.go b/server/channels/app/slashcommands/command_online.go index aaa8a9c0e2..190ff60197 100644 --- a/server/channels/app/slashcommands/command_online.go +++ b/server/channels/app/slashcommands/command_online.go @@ -34,7 +34,7 @@ func (*OnlineProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comma } } -func (*OnlineProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*OnlineProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { a.SetStatusOnline(args.UserId, true) return &model.CommandResponse{ResponseType: model.CommandResponseTypeEphemeral, Text: args.T("api.command_online.success")} diff --git a/server/channels/app/slashcommands/command_remote.go b/server/channels/app/slashcommands/command_remote.go index 0fedee496c..bfdad9efab 100644 --- a/server/channels/app/slashcommands/command_remote.go +++ b/server/channels/app/slashcommands/command_remote.go @@ -68,7 +68,7 @@ func (rp *RemoteProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Co } } -func (rp *RemoteProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (rp *RemoteProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if !a.HasPermissionTo(args.UserId, model.PermissionManageSecureConnections) { return responsef(args.T("api.command_remote.permission_required", map[string]any{"Permission": "manage_secure_connections"})) } diff --git a/server/channels/app/slashcommands/command_remove.go b/server/channels/app/slashcommands/command_remove.go index a34906e98a..8799fe48c4 100644 --- a/server/channels/app/slashcommands/command_remove.go +++ b/server/channels/app/slashcommands/command_remove.go @@ -57,15 +57,15 @@ func (*KickProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Command } } -func (*RemoveProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*RemoveProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return doCommand(a, c, args, message) } -func (*KickProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*KickProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { return doCommand(a, c, args, message) } -func doCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func doCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { channel, err := a.GetChannel(c, args.ChannelId) if err != nil { return &model.CommandResponse{ diff --git a/server/channels/app/slashcommands/command_search.go b/server/channels/app/slashcommands/command_search.go index b7b3dacfc3..259e7beac3 100644 --- a/server/channels/app/slashcommands/command_search.go +++ b/server/channels/app/slashcommands/command_search.go @@ -35,7 +35,7 @@ func (search *SearchProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *mode } } -func (search *SearchProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (search *SearchProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { // This command is handled client-side and shouldn't hit the server. return &model.CommandResponse{ Text: args.T("api.command_search.unsupported.app_error"), diff --git a/server/channels/app/slashcommands/command_settings.go b/server/channels/app/slashcommands/command_settings.go index 96071729c0..3f4048577d 100644 --- a/server/channels/app/slashcommands/command_settings.go +++ b/server/channels/app/slashcommands/command_settings.go @@ -35,7 +35,7 @@ func (settings *SettingsProvider) GetCommand(a *app.App, T i18n.TranslateFunc) * } } -func (settings *SettingsProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (settings *SettingsProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { // This command is handled client-side and shouldn't hit the server. return &model.CommandResponse{ Text: args.T("api.command_settings.unsupported.app_error"), diff --git a/server/channels/app/slashcommands/command_share.go b/server/channels/app/slashcommands/command_share.go index 1281253aab..eda2985aa0 100644 --- a/server/channels/app/slashcommands/command_share.go +++ b/server/channels/app/slashcommands/command_share.go @@ -119,7 +119,7 @@ func (sp *ShareProvider) getAutoCompleteUnInviteRemote(a *app.App, _ *model.Comm } } -func (sp *ShareProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (sp *ShareProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { if !a.HasPermissionTo(args.UserId, model.PermissionManageSharedChannels) { return responsef(args.T("api.command_share.permission_required", map[string]any{"Permission": "manage_shared_channels"})) } diff --git a/server/channels/app/slashcommands/command_shortcuts.go b/server/channels/app/slashcommands/command_shortcuts.go index 26b2b5f706..2e79d35304 100644 --- a/server/channels/app/slashcommands/command_shortcuts.go +++ b/server/channels/app/slashcommands/command_shortcuts.go @@ -35,7 +35,7 @@ func (*ShortcutsProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Co } } -func (*ShortcutsProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*ShortcutsProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { // This command is handled client-side and shouldn't hit the server. return &model.CommandResponse{ Text: args.T("api.command_shortcuts.unsupported.app_error"), diff --git a/server/channels/app/slashcommands/command_shrug.go b/server/channels/app/slashcommands/command_shrug.go index 3af4dbefb0..af7444b9ad 100644 --- a/server/channels/app/slashcommands/command_shrug.go +++ b/server/channels/app/slashcommands/command_shrug.go @@ -35,7 +35,7 @@ func (*ShrugProvider) GetCommand(a *app.App, T i18n.TranslateFunc) *model.Comman } } -func (*ShrugProvider) DoCommand(a *app.App, c request.CTX, args *model.CommandArgs, message string) *model.CommandResponse { +func (*ShrugProvider) DoCommand(a *app.App, c *request.Context, args *model.CommandArgs, message string) *model.CommandResponse { rmsg := `¯\\\_(ツ)\_/¯` if message != "" { rmsg = message + " " + rmsg diff --git a/server/channels/app/syncables.go b/server/channels/app/syncables.go index 5a060f2b97..2e1513f1fa 100644 --- a/server/channels/app/syncables.go +++ b/server/channels/app/syncables.go @@ -33,7 +33,7 @@ func (a *App) createDefaultChannelMemberships(c request.CTX, params model.Create return err } - tmem, err := a.GetTeamMember(channel.TeamId, userChannel.UserID) + tmem, err := a.GetTeamMember(c, channel.TeamId, userChannel.UserID) if err != nil && err.Id != "app.team.get_member.missing.app_error" { return err } diff --git a/server/channels/app/syncables_test.go b/server/channels/app/syncables_test.go index c1b80a90f5..e7738cf572 100644 --- a/server/channels/app/syncables_test.go +++ b/server/channels/app/syncables_test.go @@ -109,7 +109,7 @@ func TestCreateDefaultMemberships(t *testing.T) { } // Singer should be in team and channel - _, err = th.App.GetTeamMember(singersTeam.Id, singer1.Id) + _, err = th.App.GetTeamMember(th.Context, singersTeam.Id, singer1.Id) if err != nil { t.Errorf("error retrieving team member: %s", err.Error()) } @@ -137,7 +137,7 @@ func TestCreateDefaultMemberships(t *testing.T) { } // Scientist should not be in team or channel - _, err = th.App.GetTeamMember(nerdsTeam.Id, scientist1.Id) + _, err = th.App.GetTeamMember(th.Context, nerdsTeam.Id, scientist1.Id) if err.Id != "app.team.get_member.missing.app_error" { t.Errorf("wrong error: %s", err.Id) } @@ -179,7 +179,7 @@ func TestCreateDefaultMemberships(t *testing.T) { } // Scientist should be in team but not the channel - _, err = th.App.GetTeamMember(nerdsTeam.Id, scientist1.Id) + _, err = th.App.GetTeamMember(th.Context, nerdsTeam.Id, scientist1.Id) if err != nil { t.Errorf("error retrieving team member: %s", err.Error()) } @@ -247,7 +247,7 @@ func TestCreateDefaultMemberships(t *testing.T) { } // Singer should not be in team or channel - tMember, err := th.App.GetTeamMember(singersTeam.Id, singer1.Id) + tMember, err := th.App.GetTeamMember(th.Context, singersTeam.Id, singer1.Id) if err != nil { t.Errorf("error retrieving team member: %s", err.Error()) } @@ -608,7 +608,7 @@ func TestSyncSyncableRoles(t *testing.T) { require.Nil(t, err) for _, user := range []*model.User{user1, user2} { - tm, err := th.App.GetTeamMember(team.Id, user.Id) + tm, err := th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, err) require.True(t, tm.SchemeAdmin) diff --git a/server/channels/app/team.go b/server/channels/app/team.go index cdc217aeb7..3319dce698 100644 --- a/server/channels/app/team.go +++ b/server/channels/app/team.go @@ -36,8 +36,8 @@ type teamServiceWrapper struct { app AppIface } -func (w *teamServiceWrapper) GetMember(teamID, userID string) (*model.TeamMember, *model.AppError) { - return w.app.GetTeamMember(teamID, userID) +func (w *teamServiceWrapper) GetMember(c request.CTX, teamID, userID string) (*model.TeamMember, *model.AppError) { + return w.app.GetTeamMember(c, teamID, userID) } func (w *teamServiceWrapper) CreateMember(ctx *request.Context, teamID, userID string) (*model.TeamMember, *model.AppError) { @@ -420,8 +420,8 @@ func (a *App) GetSchemeRolesForTeam(teamID string) (string, string, string, *mod return model.TeamGuestRoleId, model.TeamUserRoleId, model.TeamAdminRoleId, nil } -func (a *App) UpdateTeamMemberRoles(teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) { - member, nErr := a.Srv().Store().Team().GetMember(context.Background(), teamID, userID) +func (a *App) UpdateTeamMemberRoles(c request.CTX, teamID string, userID string, newRoles string) (*model.TeamMember, *model.AppError) { + member, nErr := a.Srv().Store().Team().GetMember(c, teamID, userID) if nErr != nil { var nfErr *store.ErrNotFound switch { @@ -504,8 +504,8 @@ func (a *App) UpdateTeamMemberRoles(teamID string, userID string, newRoles strin return member, nil } -func (a *App) UpdateTeamMemberSchemeRoles(teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) { - member, err := a.GetTeamMember(teamID, userID) +func (a *App) UpdateTeamMemberSchemeRoles(c request.CTX, teamID string, userID string, isSchemeGuest bool, isSchemeUser bool, isSchemeAdmin bool) (*model.TeamMember, *model.AppError) { + member, err := a.GetTeamMember(c, teamID, userID) if err != nil { return nil, err } @@ -759,7 +759,7 @@ func (a *App) AddUserToTeamByInviteId(c *request.Context, inviteId string, userI } func (a *App) JoinUserToTeam(c request.CTX, team *model.Team, user *model.User, userRequestorId string) (*model.TeamMember, *model.AppError) { - teamMember, alreadyAdded, err := a.ch.srv.teamService.JoinUserToTeam(team, user) + teamMember, alreadyAdded, err := a.ch.srv.teamService.JoinUserToTeam(c, team, user) if err != nil { var appErr *model.AppError var conflictErr *store.ErrConflict @@ -793,7 +793,7 @@ func (a *App) JoinUserToTeam(c request.CTX, team *model.Team, user *model.User, TeamID: team.Id, ExcludeTeam: false, } - if _, err := a.createInitialSidebarCategories(user.Id, opts); err != nil { + if _, err := a.createInitialSidebarCategories(c, user.Id, opts); err != nil { mlog.Warn( "Encountered an issue creating default sidebar categories.", mlog.String("user_id", user.Id), @@ -993,8 +993,8 @@ func (a *App) GetTeamsForUser(userID string) ([]*model.Team, *model.AppError) { return teams, nil } -func (a *App) GetTeamMember(teamID, userID string) (*model.TeamMember, *model.AppError) { - teamMember, err := a.Srv().Store().Team().GetMember(sqlstore.WithMaster(context.Background()), teamID, userID) +func (a *App) GetTeamMember(c request.CTX, teamID, userID string) (*model.TeamMember, *model.AppError) { + teamMember, err := a.Srv().Store().Team().GetMember(sqlstore.RequestContextWithMaster(c), teamID, userID) if err != nil { var nfErr *store.ErrNotFound switch { @@ -1008,8 +1008,8 @@ func (a *App) GetTeamMember(teamID, userID string) (*model.TeamMember, *model.Ap return teamMember, nil } -func (a *App) GetTeamMembersForUser(userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) { - teamMembers, err := a.Srv().Store().Team().GetTeamsForUser(context.Background(), userID, excludeTeamID, includeDeleted) +func (a *App) GetTeamMembersForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, *model.AppError) { + teamMembers, err := a.Srv().Store().Team().GetTeamsForUser(c, userID, excludeTeamID, includeDeleted) if err != nil { return nil, model.NewAppError("GetTeamMembersForUser", "app.team.get_members.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } @@ -1237,7 +1237,7 @@ func (a *App) postProcessTeamMemberLeave(c request.CTX, teamMember *model.TeamMe } func (a *App) LeaveTeam(c request.CTX, team *model.Team, user *model.User, requestorId string) *model.AppError { - teamMember, err := a.GetTeamMember(team.Id, user.Id) + teamMember, err := a.GetTeamMember(c, team.Id, user.Id) if err != nil { return model.NewAppError("LeaveTeam", "api.team.remove_user_from_team.missing.app_error", nil, "", http.StatusBadRequest).Wrap(err) } diff --git a/server/channels/app/team_test.go b/server/channels/app/team_test.go index 4ebe652b26..b2ab165f3e 100644 --- a/server/channels/app/team_test.go +++ b/server/channels/app/team_test.go @@ -1027,7 +1027,7 @@ func TestLeaveTeamPanic(t *testing.T) { mockLicenseStore.On("Get", "").Return(&model.LicenseRecord{}, nil) mockTeamStore := mocks.TeamStore{} - mockTeamStore.On("GetMember", sqlstore.WithMaster(context.Background()), "myteam", "userID").Return(&model.TeamMember{TeamId: "myteam", UserId: "userID"}, nil) + mockTeamStore.On("GetMember", sqlstore.RequestContextWithMaster(th.Context), "myteam", "userID").Return(&model.TeamMember{TeamId: "myteam", UserId: "userID"}, nil) mockTeamStore.On("UpdateMember", mock.Anything).Return(nil, errors.New("repro error")) // This is the line that triggers the error mockStore.On("Channel").Return(&mockChannelStore) @@ -1297,7 +1297,7 @@ func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { _, _, err := th.App.AddUserToTeam(th.Context, th.BasicTeam.Id, ruser.Id, "") require.Nil(t, err) - _, err = th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, "team_user") + _, err = th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, "team_user") require.NotNil(t, err, "Should fail when try to modify the guest role") }) @@ -1308,7 +1308,7 @@ func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { _, _, err := th.App.AddUserToTeam(th.Context, th.BasicTeam.Id, ruser.Id, "") require.Nil(t, err) - _, err = th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, "team_guest") + _, err = th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, "team_guest") require.NotNil(t, err, "Should fail when try to modify the guest role") }) @@ -1319,7 +1319,7 @@ func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { _, _, err := th.App.AddUserToTeam(th.Context, th.BasicTeam.Id, ruser.Id, "") require.Nil(t, err) - _, err = th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, "team_user team_admin") + _, err = th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, "team_user team_admin") require.Nil(t, err, "Should work when you not modify guest role") }) @@ -1333,7 +1333,7 @@ func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { _, err = th.App.CreateRole(&model.Role{Name: "custom", DisplayName: "custom", Description: "custom"}) require.Nil(t, err) - _, err = th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, "team_guest custom") + _, err = th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, "team_guest custom") require.Nil(t, err, "Should work when you not modify guest role") }) @@ -1344,7 +1344,7 @@ func TestUpdateTeamMemberRolesChangingGuest(t *testing.T) { _, _, err := th.App.AddUserToTeam(th.Context, th.BasicTeam.Id, ruser.Id, "") require.Nil(t, err) - _, err = th.App.UpdateTeamMemberRoles(th.BasicTeam.Id, ruser.Id, "team_guest team_user") + _, err = th.App.UpdateTeamMemberRoles(th.Context, th.BasicTeam.Id, ruser.Id, "team_guest team_user") require.NotNil(t, err, "Should work when you not modify guest role") }) } diff --git a/server/channels/app/teams/teams.go b/server/channels/app/teams/teams.go index c1ef533e15..e4b959ae27 100644 --- a/server/channels/app/teams/teams.go +++ b/server/channels/app/teams/teams.go @@ -4,10 +4,9 @@ package teams import ( - "context" - "github.com/mattermost/mattermost/server/public/model" "github.com/mattermost/mattermost/server/public/shared/i18n" + "github.com/mattermost/mattermost/server/public/shared/request" ) func (ts *TeamService) CreateTeam(team *model.Team) (*model.Team, error) { @@ -130,7 +129,7 @@ func (ts *TeamService) PatchTeam(teamID string, patch *model.TeamPatch) (*model. // 1. a pointer to the team member, if successful // 2. a boolean: true if the user has a non-deleted team member for that team already, otherwise false. // 3. a pointer to an AppError if something went wrong. -func (ts *TeamService) JoinUserToTeam(team *model.Team, user *model.User) (*model.TeamMember, bool, error) { +func (ts *TeamService) JoinUserToTeam(c request.CTX, team *model.Team, user *model.User) (*model.TeamMember, bool, error) { if !ts.IsTeamEmailAllowed(user, team) { return nil, false, AcceptedDomainError } @@ -155,7 +154,7 @@ func (ts *TeamService) JoinUserToTeam(team *model.Team, user *model.User) (*mode tm.SchemeAdmin = true } - rtm, err := ts.store.GetMember(context.Background(), team.Id, user.Id) + rtm, err := ts.store.GetMember(c, team.Id, user.Id) if err != nil { // Membership appears to be missing. Lets try to add. tmr, nErr := ts.store.SaveMember(tm, *ts.config().TeamSettings.MaxUsersPerTeam) @@ -221,8 +220,8 @@ func (ts *TeamService) RemoveTeamMember(teamMember *model.TeamMember) error { } // GetMember return the team member from the team. -func (ts *TeamService) GetMember(teamID string, userID string) (*model.TeamMember, error) { - member, err := ts.store.GetMember(context.Background(), teamID, userID) +func (ts *TeamService) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) { + member, err := ts.store.GetMember(c, teamID, userID) if err != nil { return nil, err } diff --git a/server/channels/app/teams/teams_test.go b/server/channels/app/teams/teams_test.go index ab65925100..a4ce247939 100644 --- a/server/channels/app/teams/teams_test.go +++ b/server/channels/app/teams/teams_test.go @@ -59,7 +59,7 @@ func TestJoinUserToTeam(t *testing.T) { ruser := th.CreateUser(&user) defer th.DeleteUser(&user) - _, alreadyAdded, err := th.service.JoinUserToTeam(team, ruser) + _, alreadyAdded, err := th.service.JoinUserToTeam(th.Context, team, ruser) require.False(t, alreadyAdded, "Should return already added equal to false") require.NoError(t, err) }) @@ -69,10 +69,10 @@ func TestJoinUserToTeam(t *testing.T) { ruser := th.CreateUser(&user) defer th.DeleteUser(&user) - _, _, err := th.service.JoinUserToTeam(team, ruser) + _, _, err := th.service.JoinUserToTeam(th.Context, team, ruser) require.NoError(t, err) - _, alreadyAdded, err := th.service.JoinUserToTeam(team, ruser) + _, alreadyAdded, err := th.service.JoinUserToTeam(th.Context, team, ruser) require.True(t, alreadyAdded, "Should return already added") require.NoError(t, err) }) @@ -82,12 +82,12 @@ func TestJoinUserToTeam(t *testing.T) { ruser := th.CreateUser(&user) defer th.DeleteUser(&user) - member, _, err := th.service.JoinUserToTeam(team, ruser) + member, _, err := th.service.JoinUserToTeam(th.Context, team, ruser) require.NoError(t, err) err = th.service.RemoveTeamMember(member) require.NoError(t, err) - _, alreadyAdded, err := th.service.JoinUserToTeam(team, ruser) + _, alreadyAdded, err := th.service.JoinUserToTeam(th.Context, team, ruser) require.False(t, alreadyAdded, "Should return already added equal to false") require.NoError(t, err) }) @@ -101,10 +101,10 @@ func TestJoinUserToTeam(t *testing.T) { defer th.DeleteUser(&user1) defer th.DeleteUser(&user2) - _, _, err := th.service.JoinUserToTeam(team, ruser1) + _, _, err := th.service.JoinUserToTeam(th.Context, team, ruser1) require.NoError(t, err) - _, _, err = th.service.JoinUserToTeam(team, ruser2) + _, _, err = th.service.JoinUserToTeam(th.Context, team, ruser2) require.Error(t, err, "Should fail") }) @@ -118,14 +118,14 @@ func TestJoinUserToTeam(t *testing.T) { defer th.DeleteUser(&user1) defer th.DeleteUser(&user2) - member, _, err := th.service.JoinUserToTeam(team, ruser1) + member, _, err := th.service.JoinUserToTeam(th.Context, team, ruser1) require.NoError(t, err) err = th.service.RemoveTeamMember(member) require.NoError(t, err) - _, _, err = th.service.JoinUserToTeam(team, ruser2) + _, _, err = th.service.JoinUserToTeam(th.Context, team, ruser2) require.NoError(t, err) - _, _, err = th.service.JoinUserToTeam(team, ruser1) + _, _, err = th.service.JoinUserToTeam(th.Context, team, ruser1) require.Error(t, err, "Should fail") }) } diff --git a/server/channels/app/upload.go b/server/channels/app/upload.go index 7de74b999c..b0e27dc2d5 100644 --- a/server/channels/app/upload.go +++ b/server/channels/app/upload.go @@ -156,7 +156,7 @@ func (a *App) CreateUploadSession(c request.CTX, us *model.UploadSession) (*mode } func (a *App) GetUploadSession(c request.CTX, uploadId string) (*model.UploadSession, *model.AppError) { - us, err := a.Srv().Store().UploadSession().Get(c.Context(), uploadId) + us, err := a.Srv().Store().UploadSession().Get(c, uploadId) if err != nil { var nfErr *store.ErrNotFound switch { diff --git a/server/channels/app/user.go b/server/channels/app/user.go index a2c4b7e0a9..a52be3787c 100644 --- a/server/channels/app/user.go +++ b/server/channels/app/user.go @@ -917,7 +917,7 @@ func (a *App) UpdatePasswordAsUser(c request.CTX, userID, currentPassword, newPa return a.UpdatePasswordSendEmail(c, user, newPassword, T("api.user.update_password.menu")) } -func (a *App) userDeactivated(c request.CTX, userID string) *model.AppError { +func (a *App) userDeactivated(c *request.Context, userID string) *model.AppError { a.SetStatusOffline(userID, false) user, err := a.GetUser(userID) @@ -966,7 +966,7 @@ func (a *App) invalidateUserChannelMembersCaches(c request.CTX, userID string) * return nil } -func (a *App) UpdateActive(c request.CTX, user *model.User, active bool) (*model.User, *model.AppError) { +func (a *App) UpdateActive(c *request.Context, user *model.User, active bool) (*model.User, *model.AppError) { user.UpdateAt = model.GetMillis() if active { user.DeleteAt = 0 @@ -990,7 +990,7 @@ func (a *App) UpdateActive(c request.CTX, user *model.User, active bool) (*model ruser := userUpdate.New if !active { - if err := a.RevokeAllSessions(ruser.Id); err != nil { + if err := a.RevokeAllSessions(c, ruser.Id); err != nil { return nil, err } if err := a.userDeactivated(c, ruser.Id); err != nil { @@ -1023,7 +1023,7 @@ func (a *App) DeactivateGuests(c *request.Context) *model.AppError { } for _, userID := range userIDs { - if err := a.Srv().Platform().RevokeAllSessions(userID); err != nil { + if err := a.Srv().Platform().RevokeAllSessions(c, userID); err != nil { return model.NewAppError("DeactivateGuests", "app.user.update_active_for_multiple_users.updating.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } } @@ -1300,7 +1300,7 @@ func (a *App) UpdateUser(c request.CTX, user *model.User, sendNotifications bool return newUser, nil } -func (a *App) UpdateUserActive(c request.CTX, userID string, active bool) *model.AppError { +func (a *App) UpdateUserActive(c *request.Context, userID string, active bool) *model.AppError { user, err := a.GetUser(userID) if err != nil { @@ -2155,8 +2155,8 @@ func (a *App) UpdateOAuthUserAttrs(c *request.Context, userData io.Reader, user return nil } -func (a *App) RestrictUsersGetByPermissions(userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) { - restrictions, err := a.GetViewUsersRestrictions(userID) +func (a *App) RestrictUsersGetByPermissions(c request.CTX, userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) { + restrictions, err := a.GetViewUsersRestrictions(c, userID) if err != nil { return nil, err } @@ -2211,8 +2211,8 @@ func (a *App) filterNonGroupUsers(userIDs []string, groupUsers []*model.User) ([ return nonMemberIds, nil } -func (a *App) RestrictUsersSearchByPermissions(userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) { - restrictions, err := a.GetViewUsersRestrictions(userID) +func (a *App) RestrictUsersSearchByPermissions(c request.CTX, userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) { + restrictions, err := a.GetViewUsersRestrictions(c, userID) if err != nil { return nil, err } @@ -2221,12 +2221,12 @@ func (a *App) RestrictUsersSearchByPermissions(userID string, options *model.Use return options, nil } -func (a *App) UserCanSeeOtherUser(userID string, otherUserId string) (bool, *model.AppError) { +func (a *App) UserCanSeeOtherUser(c request.CTX, userID string, otherUserId string) (bool, *model.AppError) { if userID == otherUserId { return true, nil } - restrictions, err := a.GetViewUsersRestrictions(userID) + restrictions, err := a.GetViewUsersRestrictions(c, userID) if err != nil { return false, err } @@ -2267,7 +2267,7 @@ func (a *App) userBelongsToChannels(userID string, channelIDs []string) (bool, * return belongs, nil } -func (a *App) GetViewUsersRestrictions(userID string) (*model.ViewUsersRestrictions, *model.AppError) { +func (a *App) GetViewUsersRestrictions(c request.CTX, userID string) (*model.ViewUsersRestrictions, *model.AppError) { if a.HasPermissionTo(userID, model.PermissionViewMembers) { return nil, nil } @@ -2279,7 +2279,7 @@ func (a *App) GetViewUsersRestrictions(userID string) (*model.ViewUsersRestricti teamIDsWithPermission := []string{} for _, teamID := range teamIDs { - if a.HasPermissionToTeam(userID, teamID, model.PermissionViewMembers) { + if a.HasPermissionToTeam(c, userID, teamID, model.PermissionViewMembers) { teamIDsWithPermission = append(teamIDsWithPermission, teamID) } } @@ -2322,12 +2322,12 @@ func (a *App) PromoteGuestToUser(c *request.Context, user *model.User, requestor c.Logger().Warn("Failed to get user on promote guest to user", mlog.Err(err)) } else { a.sendUpdatedUserEvent(*promotedUser) - if uErr := a.ch.srv.platform.UpdateSessionsIsGuest(promotedUser.Id, promotedUser.IsGuest()); uErr != nil { + if uErr := a.ch.srv.platform.UpdateSessionsIsGuest(c, promotedUser.Id, promotedUser.IsGuest()); uErr != nil { c.Logger().Warn("Unable to update user sessions", mlog.String("user_id", promotedUser.Id), mlog.Err(uErr)) } } - teamMembers, err := a.GetTeamMembersForUser(user.Id, "", true) + teamMembers, err := a.GetTeamMembersForUser(c, user.Id, "", true) if err != nil { c.Logger().Warn("Failed to get team members for user on promote guest to user", mlog.Err(err)) } @@ -2359,7 +2359,7 @@ func (a *App) PromoteGuestToUser(c *request.Context, user *model.User, requestor // DemoteUserToGuest Convert user's roles and all his membership's roles from // regular user roles to guest roles. -func (a *App) DemoteUserToGuest(c request.CTX, user *model.User) *model.AppError { +func (a *App) DemoteUserToGuest(c *request.Context, user *model.User) *model.AppError { demotedUser, nErr := a.ch.srv.userService.DemoteUserToGuest(user) a.InvalidateCacheForUser(user.Id) if nErr != nil { @@ -2367,11 +2367,11 @@ func (a *App) DemoteUserToGuest(c request.CTX, user *model.User) *model.AppError } a.sendUpdatedUserEvent(*demotedUser) - if uErr := a.ch.srv.platform.UpdateSessionsIsGuest(demotedUser.Id, demotedUser.IsGuest()); uErr != nil { + if uErr := a.ch.srv.platform.UpdateSessionsIsGuest(c, demotedUser.Id, demotedUser.IsGuest()); uErr != nil { c.Logger().Warn("Unable to update user sessions", mlog.String("user_id", demotedUser.Id), mlog.Err(uErr)) } - teamMembers, err := a.GetTeamMembersForUser(user.Id, "", true) + teamMembers, err := a.GetTeamMembersForUser(c, user.Id, "", true) if err != nil { c.Logger().Warn("Failed to get team members for users on demote user to guest", mlog.Err(err)) } diff --git a/server/channels/app/user_test.go b/server/channels/app/user_test.go index 7732fd99dc..bf2ac7eb31 100644 --- a/server/channels/app/user_test.go +++ b/server/channels/app/user_test.go @@ -1223,7 +1223,7 @@ func TestGetViewUsersRestrictions(t *testing.T) { th.LinkUserToTeam(user1, team1) th.LinkUserToTeam(user1, team2) - th.App.UpdateTeamMemberRoles(team1.Id, user1.Id, "team_user team_admin") + th.App.UpdateTeamMemberRoles(th.Context, team1.Id, user1.Id, "team_user team_admin") team1channel1 := th.CreateChannel(th.Context, team1) team1channel2 := th.CreateChannel(th.Context, team1) @@ -1262,7 +1262,7 @@ func TestGetViewUsersRestrictions(t *testing.T) { } t.Run("VIEW_MEMBERS permission granted at system level", func(t *testing.T) { - restrictions, err := th.App.GetViewUsersRestrictions(user1.Id) + restrictions, err := th.App.GetViewUsersRestrictions(th.Context, user1.Id) require.Nil(t, err) assert.Nil(t, restrictions) @@ -1279,7 +1279,7 @@ func TestGetViewUsersRestrictions(t *testing.T) { require.Nil(t, addPermission(teamUserRole, model.PermissionViewMembers.Id)) defer removePermission(teamUserRole, model.PermissionViewMembers.Id) - restrictions, err := th.App.GetViewUsersRestrictions(user1.Id) + restrictions, err := th.App.GetViewUsersRestrictions(th.Context, user1.Id) require.Nil(t, err) assert.NotNil(t, restrictions) @@ -1295,7 +1295,7 @@ func TestGetViewUsersRestrictions(t *testing.T) { require.Nil(t, removePermission(systemUserRole, model.PermissionViewMembers.Id)) defer addPermission(systemUserRole, model.PermissionViewMembers.Id) - restrictions, err := th.App.GetViewUsersRestrictions(user1.Id) + restrictions, err := th.App.GetViewUsersRestrictions(th.Context, user1.Id) require.Nil(t, err) assert.NotNil(t, restrictions) @@ -1315,7 +1315,7 @@ func TestGetViewUsersRestrictions(t *testing.T) { require.Nil(t, addPermission(teamAdminRole, model.PermissionViewMembers.Id)) defer removePermission(teamAdminRole, model.PermissionViewMembers.Id) - restrictions, err := th.App.GetViewUsersRestrictions(user1.Id) + restrictions, err := th.App.GetViewUsersRestrictions(th.Context, user1.Id) require.Nil(t, err) assert.NotNil(t, restrictions) @@ -1355,7 +1355,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest := th.CreateGuest() require.Equal(t, "system_guest", guest.Roles) th.LinkUserToTeam(guest, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) require.Nil(t, err) require.True(t, teamMember.SchemeGuest) require.False(t, teamMember.SchemeUser) @@ -1365,7 +1365,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest, err = th.App.GetUser(guest.Id) assert.Nil(t, err) assert.Equal(t, "system_user", guest.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeGuest) assert.True(t, teamMember.SchemeUser) @@ -1375,7 +1375,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest := th.CreateGuest() require.Equal(t, "system_guest", guest.Roles) th.LinkUserToTeam(guest, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) require.Nil(t, err) require.True(t, teamMember.SchemeGuest) require.False(t, teamMember.SchemeUser) @@ -1389,7 +1389,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest, err = th.App.GetUser(guest.Id) assert.Nil(t, err) assert.Equal(t, "system_user", guest.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeGuest) assert.True(t, teamMember.SchemeUser) @@ -1403,7 +1403,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest := th.CreateGuest() require.Equal(t, "system_guest", guest.Roles) th.LinkUserToTeam(guest, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) require.Nil(t, err) require.True(t, teamMember.SchemeGuest) require.False(t, teamMember.SchemeUser) @@ -1421,7 +1421,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest, err = th.App.GetUser(guest.Id) assert.Nil(t, err) assert.Equal(t, "system_user", guest.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeGuest) assert.True(t, teamMember.SchemeUser) @@ -1439,7 +1439,7 @@ func TestPromoteGuestToUser(t *testing.T) { guest := th.CreateGuest() require.Equal(t, "system_guest", guest.Roles) th.LinkUserToTeam(guest, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, guest.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, guest.Id) require.Nil(t, err) require.True(t, teamMember.SchemeGuest) require.False(t, teamMember.SchemeUser) @@ -1470,7 +1470,7 @@ func TestDemoteUserToGuest(t *testing.T) { user := th.CreateUser() require.Equal(t, "system_user", user.Roles) th.LinkUserToTeam(user, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) require.Nil(t, err) require.True(t, teamMember.SchemeUser) require.False(t, teamMember.SchemeGuest) @@ -1518,7 +1518,7 @@ func TestDemoteUserToGuest(t *testing.T) { user := th.CreateUser() require.Equal(t, "system_user", user.Roles) th.LinkUserToTeam(user, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) require.Nil(t, err) require.True(t, teamMember.SchemeUser) require.False(t, teamMember.SchemeGuest) @@ -1528,7 +1528,7 @@ func TestDemoteUserToGuest(t *testing.T) { user, err = th.App.GetUser(user.Id) assert.Nil(t, err) assert.Equal(t, "system_guest", user.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeUser) assert.True(t, teamMember.SchemeGuest) @@ -1538,7 +1538,7 @@ func TestDemoteUserToGuest(t *testing.T) { user := th.CreateUser() require.Equal(t, "system_user", user.Roles) th.LinkUserToTeam(user, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) require.Nil(t, err) require.True(t, teamMember.SchemeUser) require.False(t, teamMember.SchemeGuest) @@ -1552,7 +1552,7 @@ func TestDemoteUserToGuest(t *testing.T) { user, err = th.App.GetUser(user.Id) assert.Nil(t, err) assert.Equal(t, "system_guest", user.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeUser) assert.True(t, teamMember.SchemeGuest) @@ -1566,7 +1566,7 @@ func TestDemoteUserToGuest(t *testing.T) { user := th.CreateUser() require.Equal(t, "system_user", user.Roles) th.LinkUserToTeam(user, th.BasicTeam) - teamMember, err := th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err := th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) require.Nil(t, err) require.True(t, teamMember.SchemeUser) require.False(t, teamMember.SchemeGuest) @@ -1584,7 +1584,7 @@ func TestDemoteUserToGuest(t *testing.T) { user, err = th.App.GetUser(user.Id) assert.Nil(t, err) assert.Equal(t, "system_guest", user.Roles) - teamMember, err = th.App.GetTeamMember(th.BasicTeam.Id, user.Id) + teamMember, err = th.App.GetTeamMember(th.Context, th.BasicTeam.Id, user.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeUser) assert.True(t, teamMember.SchemeGuest) @@ -1605,9 +1605,9 @@ func TestDemoteUserToGuest(t *testing.T) { team := th.CreateTeam() th.LinkUserToTeam(user, team) - th.App.UpdateTeamMemberRoles(team.Id, user.Id, "team_user team_admin") + th.App.UpdateTeamMemberRoles(th.Context, team.Id, user.Id, "team_user team_admin") - teamMember, err := th.App.GetTeamMember(team.Id, user.Id) + teamMember, err := th.App.GetTeamMember(th.Context, team.Id, user.Id) require.Nil(t, err) require.True(t, teamMember.SchemeUser) require.True(t, teamMember.SchemeAdmin) @@ -1631,7 +1631,7 @@ func TestDemoteUserToGuest(t *testing.T) { assert.Nil(t, err) assert.Equal(t, "system_guest", user.Roles) - teamMember, err = th.App.GetTeamMember(team.Id, user.Id) + teamMember, err = th.App.GetTeamMember(th.Context, team.Id, user.Id) assert.Nil(t, err) assert.False(t, teamMember.SchemeUser) assert.False(t, teamMember.SchemeAdmin) diff --git a/server/channels/app/web_conn_test.go b/server/channels/app/web_conn_test.go index c667f8cc6c..f2fbcd23e9 100644 --- a/server/channels/app/web_conn_test.go +++ b/server/channels/app/web_conn_test.go @@ -17,7 +17,7 @@ import ( func TestWebConnShouldSendEvent(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() - session, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Roles: th.BasicUser.GetRawRoles(), TeamMembers: []*model.TeamMember{ + session, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser.Id, Roles: th.BasicUser.GetRawRoles(), TeamMembers: []*model.TeamMember{ { UserId: th.BasicUser.Id, TeamId: th.BasicTeam.Id, @@ -39,7 +39,7 @@ func TestWebConnShouldSendEvent(t *testing.T) { basicUserWc.SetSessionToken(session.Token) basicUserWc.SetSessionExpiresAt(session.ExpiresAt) - session2, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser2.Id, Roles: th.BasicUser2.GetRawRoles(), TeamMembers: []*model.TeamMember{ + session2, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser2.Id, Roles: th.BasicUser2.GetRawRoles(), TeamMembers: []*model.TeamMember{ { UserId: th.BasicUser2.Id, TeamId: th.BasicTeam.Id, @@ -61,7 +61,7 @@ func TestWebConnShouldSendEvent(t *testing.T) { basicUser2Wc.SetSessionToken(session2.Token) basicUser2Wc.SetSessionExpiresAt(session2.ExpiresAt) - session3, err := th.App.CreateSession(&model.Session{UserId: th.SystemAdminUser.Id, Roles: th.SystemAdminUser.GetRawRoles()}) + session3, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.SystemAdminUser.Id, Roles: th.SystemAdminUser.GetRawRoles()}) require.Nil(t, err) adminUserWc := &platform.WebConn{ @@ -77,7 +77,7 @@ func TestWebConnShouldSendEvent(t *testing.T) { adminUserWc.SetSessionToken(session3.Token) adminUserWc.SetSessionExpiresAt(session3.ExpiresAt) - session4, err := th.App.CreateSession(&model.Session{UserId: th.BasicUser.Id, Roles: th.BasicUser.GetRawRoles(), TeamMembers: []*model.TeamMember{ + session4, err := th.App.CreateSession(th.Context, &model.Session{UserId: th.BasicUser.Id, Roles: th.BasicUser.GetRawRoles(), TeamMembers: []*model.TeamMember{ { UserId: th.BasicUser.Id, TeamId: th.BasicTeam.Id, diff --git a/server/channels/product/api.go b/server/channels/product/api.go index 60f24927eb..2614f4f022 100644 --- a/server/channels/product/api.go +++ b/server/channels/product/api.go @@ -45,8 +45,8 @@ type PostService interface { // The service shall be registered via app.PermissionKey service key. type PermissionService interface { HasPermissionTo(userID string, permission *model.Permission) bool - HasPermissionToTeam(userID, teamID string, permission *model.Permission) bool - HasPermissionToChannel(askingUserID string, channelID string, permission *model.Permission) bool + HasPermissionToTeam(c *request.Context, userID, teamID string, permission *model.Permission) bool + HasPermissionToChannel(c *request.Context, askingUserID string, channelID string, permission *model.Permission) bool RolesGrantPermission(roleNames []string, permissionID string) bool } @@ -107,7 +107,7 @@ type UserService interface { // // The service shall be registered via app.TeamKey service key. type TeamService interface { - GetMember(teamID, userID string) (*model.TeamMember, *model.AppError) + GetMember(c request.CTX, teamID, userID string) (*model.TeamMember, *model.AppError) CreateMember(ctx *request.Context, teamID, userID string) (*model.TeamMember, *model.AppError) GetGroup(groupId string) (*model.Group, *model.AppError) GetTeam(teamID string) (*model.Team, *model.AppError) diff --git a/server/channels/store/opentracinglayer/opentracinglayer.go b/server/channels/store/opentracinglayer/opentracinglayer.go index 9e4f089f25..49355a569d 100644 --- a/server/channels/store/opentracinglayer/opentracinglayer.go +++ b/server/channels/store/opentracinglayer/opentracinglayer.go @@ -811,7 +811,7 @@ func (s *OpenTracingLayerChannelStore) CreateDirectChannel(userID *model.User, o return result, err } -func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { +func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ChannelStore.CreateInitialSidebarCategories") s.Root.Store.SetContext(newCtx) @@ -820,7 +820,7 @@ func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(userID str }() defer span.Finish() - result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts) + result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -3209,7 +3209,7 @@ func (s *OpenTracingLayerComplianceStore) GetAll(offset int, limit int) (model.C return result, err } -func (s *OpenTracingLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { +func (s *OpenTracingLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ComplianceStore.MessageExport") s.Root.Store.SetContext(newCtx) @@ -3218,7 +3218,7 @@ func (s *OpenTracingLayerComplianceStore) MessageExport(ctx context.Context, cur }() defer span.Finish() - result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit) + result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -3443,7 +3443,7 @@ func (s *OpenTracingLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64) return err } -func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { +func (s *OpenTracingLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.Get") s.Root.Store.SetContext(newCtx) @@ -3452,7 +3452,7 @@ func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCa }() defer span.Finish() - result, err := s.EmojiStore.Get(ctx, id, allowFromCache) + result, err := s.EmojiStore.Get(c, id, allowFromCache) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -3461,7 +3461,7 @@ func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCa return result, err } -func (s *OpenTracingLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { +func (s *OpenTracingLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.GetByName") s.Root.Store.SetContext(newCtx) @@ -3470,7 +3470,7 @@ func (s *OpenTracingLayerEmojiStore) GetByName(ctx request.CTX, name string, all }() defer span.Finish() - result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache) + result, err := s.EmojiStore.GetByName(c, name, allowFromCache) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -3497,7 +3497,7 @@ func (s *OpenTracingLayerEmojiStore) GetList(offset int, limit int, sort string) return result, err } -func (s *OpenTracingLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) { +func (s *OpenTracingLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.GetMultipleByName") s.Root.Store.SetContext(newCtx) @@ -3506,7 +3506,7 @@ func (s *OpenTracingLayerEmojiStore) GetMultipleByName(ctx request.CTX, names [] }() defer span.Finish() - result, err := s.EmojiStore.GetMultipleByName(ctx, names) + result, err := s.EmojiStore.GetMultipleByName(c, names) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -5197,7 +5197,7 @@ func (s *OpenTracingLayerJobStore) UpdateStatusOptimistically(id string, current return result, err } -func (s *OpenTracingLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) { +func (s *OpenTracingLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "LicenseStore.Get") s.Root.Store.SetContext(newCtx) @@ -5206,7 +5206,7 @@ func (s *OpenTracingLayerLicenseStore) Get(ctx context.Context, id string) (*mod }() defer span.Finish() - result, err := s.LicenseStore.Get(ctx, id) + result, err := s.LicenseStore.Get(c, id) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -8281,7 +8281,7 @@ func (s *OpenTracingLayerSessionStore) Cleanup(expiryTime int64, batchSize int64 return err } -func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) { +func (s *OpenTracingLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.Get") s.Root.Store.SetContext(newCtx) @@ -8290,7 +8290,7 @@ func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken }() defer span.Finish() - result, err := s.SessionStore.Get(ctx, sessionIDOrToken) + result, err := s.SessionStore.Get(c, sessionIDOrToken) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -8299,7 +8299,7 @@ func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken return result, err } -func (s *OpenTracingLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) { +func (s *OpenTracingLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.GetSessions") s.Root.Store.SetContext(newCtx) @@ -8308,7 +8308,7 @@ func (s *OpenTracingLayerSessionStore) GetSessions(userID string) ([]*model.Sess }() defer span.Finish() - result, err := s.SessionStore.GetSessions(userID) + result, err := s.SessionStore.GetSessions(c, userID) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -8407,7 +8407,7 @@ func (s *OpenTracingLayerSessionStore) RemoveAllSessions() error { return err } -func (s *OpenTracingLayerSessionStore) Save(session *model.Session) (*model.Session, error) { +func (s *OpenTracingLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.Save") s.Root.Store.SetContext(newCtx) @@ -8416,7 +8416,7 @@ func (s *OpenTracingLayerSessionStore) Save(session *model.Session) (*model.Sess }() defer span.Finish() - result, err := s.SessionStore.Save(session) + result, err := s.SessionStore.Save(c, session) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -9626,7 +9626,7 @@ func (s *OpenTracingLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) return result, err } -func (s *OpenTracingLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { +func (s *OpenTracingLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetMember") s.Root.Store.SetContext(newCtx) @@ -9635,7 +9635,7 @@ func (s *OpenTracingLayerTeamStore) GetMember(ctx context.Context, teamID string }() defer span.Finish() - result, err := s.TeamStore.GetMember(ctx, teamID, userID) + result, err := s.TeamStore.GetMember(c, teamID, userID) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -9734,7 +9734,7 @@ func (s *OpenTracingLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Te return result, err } -func (s *OpenTracingLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { +func (s *OpenTracingLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetTeamsForUser") s.Root.Store.SetContext(newCtx) @@ -9743,7 +9743,7 @@ func (s *OpenTracingLayerTeamStore) GetTeamsForUser(ctx context.Context, userID }() defer span.Finish() - result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted) + result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -10840,7 +10840,7 @@ func (s *OpenTracingLayerUploadSessionStore) Delete(id string) error { return err } -func (s *OpenTracingLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { +func (s *OpenTracingLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UploadSessionStore.Get") s.Root.Store.SetContext(newCtx) @@ -10849,7 +10849,7 @@ func (s *OpenTracingLayerUploadSessionStore) Get(ctx context.Context, id string) }() defer span.Finish() - result, err := s.UploadSessionStore.Get(ctx, id) + result, err := s.UploadSessionStore.Get(c, id) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) diff --git a/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index 6811f4c00e..9f58341f55 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -871,11 +871,11 @@ func (s *RetryLayerChannelStore) CreateDirectChannel(userID *model.User, otherUs } -func (s *RetryLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { +func (s *RetryLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { tries := 0 for { - result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts) + result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts) if err == nil { return result, nil } @@ -3577,11 +3577,11 @@ func (s *RetryLayerComplianceStore) GetAll(offset int, limit int) (model.Complia } -func (s *RetryLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { +func (s *RetryLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { tries := 0 for { - result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit) + result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit) if err == nil { return result, resultVar1, nil } @@ -3850,11 +3850,11 @@ func (s *RetryLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64) error } -func (s *RetryLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { +func (s *RetryLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { tries := 0 for { - result, err := s.EmojiStore.Get(ctx, id, allowFromCache) + result, err := s.EmojiStore.Get(c, id, allowFromCache) if err == nil { return result, nil } @@ -3871,11 +3871,11 @@ func (s *RetryLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bo } -func (s *RetryLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { +func (s *RetryLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { tries := 0 for { - result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache) + result, err := s.EmojiStore.GetByName(c, name, allowFromCache) if err == nil { return result, nil } @@ -3913,11 +3913,11 @@ func (s *RetryLayerEmojiStore) GetList(offset int, limit int, sort string) ([]*m } -func (s *RetryLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) { +func (s *RetryLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) { tries := 0 for { - result, err := s.EmojiStore.GetMultipleByName(ctx, names) + result, err := s.EmojiStore.GetMultipleByName(c, names) if err == nil { return result, nil } @@ -5878,11 +5878,11 @@ func (s *RetryLayerJobStore) UpdateStatusOptimistically(id string, currentStatus } -func (s *RetryLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) { +func (s *RetryLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) { tries := 0 for { - result, err := s.LicenseStore.Get(ctx, id) + result, err := s.LicenseStore.Get(c, id) if err == nil { return result, nil } @@ -9430,11 +9430,11 @@ func (s *RetryLayerSessionStore) Cleanup(expiryTime int64, batchSize int64) erro } -func (s *RetryLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) { +func (s *RetryLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) { tries := 0 for { - result, err := s.SessionStore.Get(ctx, sessionIDOrToken) + result, err := s.SessionStore.Get(c, sessionIDOrToken) if err == nil { return result, nil } @@ -9451,11 +9451,11 @@ func (s *RetryLayerSessionStore) Get(ctx context.Context, sessionIDOrToken strin } -func (s *RetryLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) { +func (s *RetryLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) { tries := 0 for { - result, err := s.SessionStore.GetSessions(userID) + result, err := s.SessionStore.GetSessions(c, userID) if err == nil { return result, nil } @@ -9577,11 +9577,11 @@ func (s *RetryLayerSessionStore) RemoveAllSessions() error { } -func (s *RetryLayerSessionStore) Save(session *model.Session) (*model.Session, error) { +func (s *RetryLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) { tries := 0 for { - result, err := s.SessionStore.Save(session) + result, err := s.SessionStore.Save(c, session) if err == nil { return result, nil } @@ -10990,11 +10990,11 @@ func (s *RetryLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) { } -func (s *RetryLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { +func (s *RetryLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) { tries := 0 for { - result, err := s.TeamStore.GetMember(ctx, teamID, userID) + result, err := s.TeamStore.GetMember(c, teamID, userID) if err == nil { return result, nil } @@ -11116,11 +11116,11 @@ func (s *RetryLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Team, er } -func (s *RetryLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { +func (s *RetryLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { tries := 0 for { - result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted) + result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted) if err == nil { return result, nil } @@ -12388,11 +12388,11 @@ func (s *RetryLayerUploadSessionStore) Delete(id string) error { } -func (s *RetryLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { +func (s *RetryLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) { tries := 0 for { - result, err := s.UploadSessionStore.Get(ctx, id) + result, err := s.UploadSessionStore.Get(c, id) if err == nil { return result, nil } diff --git a/server/channels/store/sqlstore/channel_store_categories.go b/server/channels/store/sqlstore/channel_store_categories.go index 7b4e55d4d7..ec2fec5522 100644 --- a/server/channels/store/sqlstore/channel_store_categories.go +++ b/server/channels/store/sqlstore/channel_store_categories.go @@ -4,13 +4,13 @@ package sqlstore import ( - "context" "fmt" sq "github.com/mattermost/squirrel" "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -20,14 +20,14 @@ type dbSelecter interface { Select(i any, query string, args ...any) error } -func (s SqlChannelStore) CreateInitialSidebarCategories(userId string, opts *store.SidebarCategorySearchOpts) (_ *model.OrderedSidebarCategories, err error) { +func (s SqlChannelStore) CreateInitialSidebarCategories(c request.CTX, userId string, opts *store.SidebarCategorySearchOpts) (_ *model.OrderedSidebarCategories, err error) { transaction, err := s.GetMasterX().Beginx() if err != nil { return nil, errors.Wrap(err, "CreateInitialSidebarCategories: begin_transaction") } defer finalizeTransactionX(transaction, &err) - teamsWithExclude, err := s.SqlStore.stores.team.GetTeamsForUser(context.Background(), userId, opts.TeamID, false) + teamsWithExclude, err := s.SqlStore.stores.team.GetTeamsForUser(c, userId, opts.TeamID, false) if err != nil { return nil, errors.Wrap(err, "CreateInitialSidebarCategories: GetTeamsForUser") } diff --git a/server/channels/store/sqlstore/compliance_store.go b/server/channels/store/sqlstore/compliance_store.go index 356fd3bfc4..5abd481942 100644 --- a/server/channels/store/sqlstore/compliance_store.go +++ b/server/channels/store/sqlstore/compliance_store.go @@ -4,7 +4,6 @@ package sqlstore import ( - "context" "database/sql" "fmt" "strings" @@ -13,6 +12,7 @@ import ( "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -271,7 +271,7 @@ func (s SqlComplianceStore) ComplianceExport(job *model.Compliance, cursor model return append(channelPosts, directMessagePosts...), cursor, nil } -func (s SqlComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { +func (s SqlComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { var args []any args = append(args, model.ChannelTypeDirect, model.ChannelTypeGroup, cursor.LastPostUpdateAt, cursor.LastPostUpdateAt, cursor.LastPostId, limit) query := @@ -318,7 +318,7 @@ func (s SqlComplianceStore) MessageExport(ctx context.Context, cursor model.Mess LIMIT ?` cposts := []*model.MessageExport{} - if err := s.GetReplicaX().SelectCtx(ctx, &cposts, query, args...); err != nil { + if err := s.GetReplicaX().SelectCtx(c.Context(), &cposts, query, args...); err != nil { return nil, cursor, errors.Wrap(err, "unable to export messages") } if len(cposts) > 0 { diff --git a/server/channels/store/sqlstore/context.go b/server/channels/store/sqlstore/context.go index 52f9cc2df3..8b43c5e415 100644 --- a/server/channels/store/sqlstore/context.go +++ b/server/channels/store/sqlstore/context.go @@ -35,8 +35,8 @@ func RequestContextWithMaster(c request.CTX) request.CTX { return c } -// hasMaster is a helper function to check whether master DB should be selected or not. -func hasMaster(ctx context.Context) bool { +// HasMaster is a helper function to check whether master DB should be selected or not. +func HasMaster(ctx context.Context) bool { if v := ctx.Value(storeContextKey(useMaster)); v != nil { if res, ok := v.(bool); ok && res { return true @@ -47,7 +47,7 @@ func hasMaster(ctx context.Context) bool { // DBXFromContext is a helper utility that returns the sqlx DB handle from a given context. func (ss *SqlStore) DBXFromContext(ctx context.Context) *sqlxDBWrapper { - if hasMaster(ctx) { + if HasMaster(ctx) { return ss.GetMasterX() } return ss.GetReplicaX() diff --git a/server/channels/store/sqlstore/context_test.go b/server/channels/store/sqlstore/context_test.go index 66c29cd59d..43d3332578 100644 --- a/server/channels/store/sqlstore/context_test.go +++ b/server/channels/store/sqlstore/context_test.go @@ -14,5 +14,5 @@ func TestContextMaster(t *testing.T) { ctx := context.Background() m := WithMaster(ctx) - assert.True(t, hasMaster(m)) + assert.True(t, HasMaster(m)) } diff --git a/server/channels/store/sqlstore/integrity_test.go b/server/channels/store/sqlstore/integrity_test.go index 26c5b07257..fce664ae47 100644 --- a/server/channels/store/sqlstore/integrity_test.go +++ b/server/channels/store/sqlstore/integrity_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -316,10 +317,10 @@ func createScheme(ss store.Store) *model.Scheme { return s } -func createSession(ss store.Store, userId string) *model.Session { +func createSession(c *request.Context, ss store.Store, userId string) *model.Session { m := model.Session{} m.UserId = userId - s, _ := ss.Session().Save(&m) + s, _ := ss.Session().Save(c, &m) return s } @@ -762,6 +763,7 @@ func TestCheckSchemesTeamsIntegrity(t *testing.T) { func TestCheckSessionsAuditsIntegrity(t *testing.T) { StoreTest(t, func(t *testing.T, ss store.Store) { + c := request.TestContext(t) store := ss.(*SqlStore) dbmap := store.GetMasterX() @@ -774,7 +776,7 @@ func TestCheckSessionsAuditsIntegrity(t *testing.T) { t.Run("should generate a report with one record", func(t *testing.T) { userId := model.NewId() - session := createSession(ss, model.NewId()) + session := createSession(c, ss, model.NewId()) sessionId := session.Id audit := createAudit(ss, userId, sessionId) dbmap.Exec(`DELETE FROM Sessions WHERE Id=?`, session.Id) @@ -1492,6 +1494,7 @@ func TestCheckUsersReactionsIntegrity(t *testing.T) { func TestCheckUsersSessionsIntegrity(t *testing.T) { StoreTest(t, func(t *testing.T, ss store.Store) { + c := request.TestContext(t) store := ss.(*SqlStore) dbmap := store.GetMasterX() @@ -1504,7 +1507,7 @@ func TestCheckUsersSessionsIntegrity(t *testing.T) { t.Run("should generate a report with one record", func(t *testing.T) { userId := model.NewId() - session := createSession(ss, userId) + session := createSession(c, ss, userId) result := checkUsersSessionsIntegrity(store) require.NoError(t, result.Err) data := result.Data.(model.RelationalIntegrityCheckData) diff --git a/server/channels/store/sqlstore/license_store.go b/server/channels/store/sqlstore/license_store.go index 8c62d6f06e..4ffd75286d 100644 --- a/server/channels/store/sqlstore/license_store.go +++ b/server/channels/store/sqlstore/license_store.go @@ -4,12 +4,11 @@ package sqlstore import ( - "context" - sq "github.com/mattermost/squirrel" "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -58,7 +57,7 @@ func (ls SqlLicenseStore) Save(license *model.LicenseRecord) error { // Get obtains the license with the provided id parameter from the database. // If the license doesn't exist it returns a model.AppError with // http.StatusNotFound in the StatusCode field. -func (ls SqlLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) { +func (ls SqlLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) { query := ls.getQueryBuilder(). Select("Id, CreateAt, Bytes"). From("Licenses"). @@ -70,7 +69,7 @@ func (ls SqlLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRec } license := &model.LicenseRecord{} - if err := ls.DBXFromContext(ctx).Get(license, queryString, args...); err != nil { + if err := ls.DBXFromContext(c.Context()).Get(license, queryString, args...); err != nil { return nil, store.NewErrNotFound("License", id) } return license, nil diff --git a/server/channels/store/sqlstore/session_store.go b/server/channels/store/sqlstore/session_store.go index be99f1282f..9849e01216 100644 --- a/server/channels/store/sqlstore/session_store.go +++ b/server/channels/store/sqlstore/session_store.go @@ -4,7 +4,6 @@ package sqlstore import ( - "context" "encoding/json" "fmt" "time" @@ -13,6 +12,7 @@ import ( "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -28,7 +28,7 @@ func newSqlSessionStore(sqlStore *SqlStore) store.SessionStore { return &SqlSessionStore{sqlStore} } -func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) { +func (me SqlSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) { if session.Id != "" { return nil, store.NewErrInvalidInput("Session", "id", session.Id) } @@ -59,7 +59,7 @@ func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) { return nil, errors.Wrapf(err, "failed to save Session with id=%s", session.Id) } - teamMembers, err := me.Team().GetTeamsForUser(context.Background(), session.UserId, "", true) + teamMembers, err := me.Team().GetTeamsForUser(c, session.UserId, "", true) if err != nil { return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId) } @@ -74,10 +74,10 @@ func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) { return session, nil } -func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*model.Session, error) { +func (me SqlSessionStore) Get(c request.CTX, sessionIdOrToken string) (*model.Session, error) { sessions := []*model.Session{} - if err := me.DBXFromContext(ctx).Select(&sessions, "SELECT * FROM Sessions WHERE Token = ? OR Id = ? LIMIT 1", sessionIdOrToken, sessionIdOrToken); err != nil { + if err := me.DBXFromContext(c.Context()).Select(&sessions, "SELECT * FROM Sessions WHERE Token = ? OR Id = ? LIMIT 1", sessionIdOrToken, sessionIdOrToken); err != nil { return nil, errors.Wrapf(err, "failed to find Sessions with sessionIdOrToken=%s", sessionIdOrToken) } if len(sessions) == 0 { @@ -86,7 +86,7 @@ func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*mo session := sessions[0] tempMembers, err := me.Team().GetTeamsForUser( - WithMaster(context.Background()), + RequestContextWithMaster(c), session.UserId, "", true) if err != nil { return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId) @@ -100,14 +100,14 @@ func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*mo return session, nil } -func (me SqlSessionStore) GetSessions(userId string) ([]*model.Session, error) { +func (me SqlSessionStore) GetSessions(c *request.Context, userId string) ([]*model.Session, error) { sessions := []*model.Session{} if err := me.GetReplicaX().Select(&sessions, "SELECT * FROM Sessions WHERE UserId = ? ORDER BY LastActivityAt DESC", userId); err != nil { return nil, errors.Wrapf(err, "failed to find Sessions with userId=%s", userId) } - teamMembers, err := me.Team().GetTeamsForUser(context.Background(), userId, "", true) + teamMembers, err := me.Team().GetTeamsForUser(c, userId, "", true) if err != nil { return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", userId) } diff --git a/server/channels/store/sqlstore/team_store.go b/server/channels/store/sqlstore/team_store.go index af2d62974a..ad50808e34 100644 --- a/server/channels/store/sqlstore/team_store.go +++ b/server/channels/store/sqlstore/team_store.go @@ -4,7 +4,6 @@ package sqlstore import ( - "context" "database/sql" "fmt" "strings" @@ -13,6 +12,7 @@ import ( "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" "github.com/mattermost/mattermost/server/v8/channels/utils" ) @@ -945,7 +945,7 @@ func (s SqlTeamStore) UpdateMember(member *model.TeamMember) (*model.TeamMember, } // GetMember returns a single member of the team that matches the teamId and userId provided as parameters. -func (s SqlTeamStore) GetMember(ctx context.Context, teamId string, userId string) (*model.TeamMember, error) { +func (s SqlTeamStore) GetMember(ctx request.CTX, teamId string, userId string) (*model.TeamMember, error) { query := s.getTeamMembersWithSchemeSelectQuery(). Where(sq.Eq{"TeamMembers.TeamId": teamId}). Where(sq.Eq{"TeamMembers.UserId": userId}) @@ -956,7 +956,7 @@ func (s SqlTeamStore) GetMember(ctx context.Context, teamId string, userId strin } var dbMember teamMemberWithSchemeRoles - err = s.DBXFromContext(ctx).Get(&dbMember, queryString, args...) + err = s.DBXFromContext(ctx.Context()).Get(&dbMember, queryString, args...) if err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("TeamMember", fmt.Sprintf("teamId=%s, userId=%s", teamId, userId)) @@ -1092,7 +1092,7 @@ func (s SqlTeamStore) GetMembersByIds(teamId string, userIds []string, restricti } // GetTeamsForUser returns a list of teams that the user is a member of. Expects userId to be passed as a parameter. It can also negative the teamID passed. -func (s SqlTeamStore) GetTeamsForUser(ctx context.Context, userId, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { +func (s SqlTeamStore) GetTeamsForUser(ctx request.CTX, userId, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { query := s.getTeamMembersWithSchemeSelectQuery(). Where(sq.Eq{"TeamMembers.UserId": userId}) @@ -1110,7 +1110,7 @@ func (s SqlTeamStore) GetTeamsForUser(ctx context.Context, userId, excludeTeamID } dbMembers := teamMemberWithSchemeRolesList{} - err = s.SqlStore.DBXFromContext(ctx).Select(&dbMembers, queryString, args...) + err = s.SqlStore.DBXFromContext(ctx.Context()).Select(&dbMembers, queryString, args...) if err != nil { return nil, errors.Wrapf(err, "failed to find TeamMembers with userId=%s", userId) } diff --git a/server/channels/store/sqlstore/upload_session_store.go b/server/channels/store/sqlstore/upload_session_store.go index 0c9e643861..f814374d46 100644 --- a/server/channels/store/sqlstore/upload_session_store.go +++ b/server/channels/store/sqlstore/upload_session_store.go @@ -4,13 +4,13 @@ package sqlstore import ( - "context" "database/sql" sq "github.com/mattermost/squirrel" "github.com/pkg/errors" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -79,7 +79,7 @@ func (us SqlUploadSessionStore) Update(session *model.UploadSession) error { return nil } -func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { +func (us SqlUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) { if !model.IsValidId(id) { return nil, errors.New("SqlUploadSessionStore.Get: id is not valid") } @@ -92,7 +92,7 @@ func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.Uplo return nil, errors.Wrap(err, "SqlUploadSessionStore.Get: failed to build query") } var session model.UploadSession - if err := us.DBXFromContext(ctx).Get(&session, query, args...); err != nil { + if err := us.DBXFromContext(c.Context()).Get(&session, query, args...); err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("UploadSession", id) } diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 0c66704e0e..67621de6e5 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -138,12 +138,12 @@ type TeamStore interface { SaveMember(member *model.TeamMember, maxUsersPerTeam int) (*model.TeamMember, error) UpdateMember(member *model.TeamMember) (*model.TeamMember, error) UpdateMultipleMembers(members []*model.TeamMember) ([]*model.TeamMember, error) - GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) + GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) GetMembers(teamID string, offset int, limit int, teamMembersGetOptions *model.TeamMembersGetOptions) ([]*model.TeamMember, error) GetMembersByIds(teamID string, userIds []string, restrictions *model.ViewUsersRestrictions) ([]*model.TeamMember, error) GetTotalMemberCount(teamID string, restrictions *model.ViewUsersRestrictions) (int64, error) GetActiveMemberCount(teamID string, restrictions *model.ViewUsersRestrictions) (int64, error) - GetTeamsForUser(ctx context.Context, userID, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) + GetTeamsForUser(c request.CTX, userID, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) GetTeamsForUserWithPagination(userID string, page, perPage int) ([]*model.TeamMember, error) GetChannelUnreadsForAllTeams(excludeTeamID, userID string) ([]*model.ChannelUnread, error) GetChannelUnreadsForTeam(teamID, userID string) ([]*model.ChannelUnread, error) @@ -278,7 +278,7 @@ type ChannelStore interface { MigrateChannelMembers(fromChannelID string, fromUserID string) (map[string]string, error) ResetAllChannelSchemes() error ClearAllCustomRoleAssignments() error - CreateInitialSidebarCategories(userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) + CreateInitialSidebarCategories(c request.CTX, userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) GetSidebarCategoriesForTeamForUser(userID, teamID string) (*model.OrderedSidebarCategories, error) GetSidebarCategories(userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) GetSidebarCategory(categoryID string) (*model.SidebarCategoryWithChannels, error) @@ -488,9 +488,9 @@ type BotStore interface { } type SessionStore interface { - Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) - Save(session *model.Session) (*model.Session, error) - GetSessions(userID string) ([]*model.Session, error) + Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) + Save(c request.CTX, session *model.Session) (*model.Session, error) + GetSessions(c *request.Context, userID string) ([]*model.Session, error) GetSessionsWithActiveDeviceIds(userID string) ([]*model.Session, error) GetSessionsExpired(thresholdMillis int64, mobileOnly bool, unnotifiedOnly bool) ([]*model.Session, error) UpdateExpiredNotify(sessionid string, notified bool) error @@ -537,7 +537,7 @@ type ComplianceStore interface { Get(id string) (*model.Compliance, error) GetAll(offset, limit int) (model.Compliances, error) ComplianceExport(compliance *model.Compliance, cursor model.ComplianceExportCursor, limit int) ([]*model.CompliancePost, model.ComplianceExportCursor, error) - MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) + MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) } type OAuthStore interface { @@ -642,7 +642,7 @@ type PreferenceStore interface { type LicenseStore interface { Save(license *model.LicenseRecord) error - Get(ctx context.Context, id string) (*model.LicenseRecord, error) + Get(c request.CTX, id string) (*model.LicenseRecord, error) GetAll() ([]*model.LicenseRecord, error) } @@ -665,9 +665,9 @@ type DesktopTokensStore interface { type EmojiStore interface { Save(emoji *model.Emoji) (*model.Emoji, error) - Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) - GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) - GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) + Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) + GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) + GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) GetList(offset, limit int, sort string) ([]*model.Emoji, error) Delete(emoji *model.Emoji, timestamp int64) error Search(name string, prefixOnly bool, limit int) ([]*model.Emoji, error) @@ -712,7 +712,7 @@ type FileInfoStore interface { type UploadSessionStore interface { Save(session *model.UploadSession) (*model.UploadSession, error) Update(session *model.UploadSession) error - Get(ctx context.Context, id string) (*model.UploadSession, error) + Get(c request.CTX, id string) (*model.UploadSession, error) GetForUser(userID string) ([]*model.UploadSession, error) Delete(id string) error } diff --git a/server/channels/store/storetest/channel_store_categories.go b/server/channels/store/storetest/channel_store_categories.go index d8d226f7fe..ec529fd420 100644 --- a/server/channels/store/storetest/channel_store_categories.go +++ b/server/channels/store/storetest/channel_store_categories.go @@ -13,6 +13,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -53,6 +54,8 @@ func setupTeam(t *testing.T, ss store.Store, userIds ...string) *model.Team { } func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { + c := request.TestContext(t) + t.Run("should create initial favorites/channels/DMs categories", func(t *testing.T) { userId := model.NewId() @@ -63,7 +66,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) assert.NoError(t, nErr) require.Len(t, res.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type) @@ -85,11 +88,11 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId2, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId2, opts) assert.NoError(t, nErr) assert.Len(t, res.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type) @@ -111,7 +114,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -119,7 +122,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team2.Id, ExcludeTeam: false, } - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts) assert.NoError(t, nErr) assert.Len(t, res.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type) @@ -140,7 +143,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -149,7 +152,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { require.Equal(t, res, initialCategories) // Calling CreateInitialSidebarCategories a second time shouldn't create any new categories - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts) assert.NoError(t, nErr) assert.NotEmpty(t, res) @@ -175,7 +178,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - _, _ = ss.Channel().CreateInitialSidebarCategories(userId, opts) + _, _ = ss.Channel().CreateInitialSidebarCategories(c, userId, opts) }() } @@ -233,7 +236,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.Len(t, categories.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type) @@ -302,7 +305,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.Len(t, categories.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type) @@ -370,7 +373,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.Len(t, categories.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type) @@ -419,7 +422,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.Len(t, categories.Categories, 3) assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type) @@ -468,7 +471,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { TeamID: t1.Id, ExcludeTeam: true, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -479,6 +482,8 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) { } func testCreateSidebarCategory(t *testing.T, ss store.Store) { + c := request.TestContext(t) + t.Run("Creating category without initial categories should fail", func(t *testing.T) { userId := model.NewId() teamId := model.NewId() @@ -504,7 +509,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -534,7 +539,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -575,7 +580,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -617,7 +622,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -682,7 +687,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) // Create the category @@ -707,6 +712,8 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) { } func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { + c := request.TestContext(t) + t.Run("should return a custom category with its Channels field set", func(t *testing.T) { userId := model.NewId() team := setupTeam(t, ss, userId) @@ -719,7 +726,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -753,7 +760,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -821,7 +828,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -864,7 +871,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { ExcludeTeam: false, } // Create the initial categories and find the channels category - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -931,7 +938,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -976,7 +983,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1018,7 +1025,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1052,7 +1059,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { TeamID: otherTeam.Id, ExcludeTeam: false, } - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1075,6 +1082,8 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { } func testGetSidebarCategories(t *testing.T, ss store.Store) { + c := request.TestContext(t) + t.Run("should return channels in the same order between different ways of getting categories", func(t *testing.T) { userId := model.NewId() team := setupTeam(t, ss, userId) @@ -1083,7 +1092,7 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1140,7 +1149,7 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) { } for _, id := range teamIds { - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, &store.SidebarCategorySearchOpts{TeamID: id}) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, &store.SidebarCategorySearchOpts{TeamID: id}) require.NoError(t, nErr) require.NotEmpty(t, res) } @@ -1176,6 +1185,8 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) { } func testUpdateSidebarCategories(t *testing.T, ss store.Store) { + c := request.TestContext(t) + t.Run("ensure the query to update SidebarCategories hasn't been polluted by UpdateSidebarCategoryOrder", func(t *testing.T) { userId := model.NewId() team := setupTeam(t, ss, userId) @@ -1185,7 +1196,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, err) require.NotEmpty(t, res) @@ -1223,7 +1234,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, err) require.NotEmpty(t, res) @@ -1254,7 +1265,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1325,7 +1336,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1390,7 +1401,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1461,7 +1472,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1475,7 +1486,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team2.Id, ExcludeTeam: false, } - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1570,7 +1581,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1583,7 +1594,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { require.Equal(t, model.SidebarCategoryChannels, channelsCategory.Type) // Create the other users' categories - res, nErr = ss.Channel().CreateInitialSidebarCategories(userId2, opts) + res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId2, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1743,7 +1754,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1802,7 +1813,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1894,7 +1905,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -1962,7 +1973,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -2017,6 +2028,8 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) { } func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string) { + c := request.TestContext(t) + userId := model.NewId() team := setupTeam(t, ss, userId) @@ -2024,7 +2037,7 @@ func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string TeamID: team.Id, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -2036,6 +2049,8 @@ func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string } func testClearSidebarOnTeamLeave(t *testing.T, ss store.Store, s SqlStore) { + c := request.TestContext(t) + t.Run("should delete all sidebar categories and channels on the team", func(t *testing.T) { userId, teamId := setupInitialSidebarCategories(t, ss) @@ -2151,7 +2166,7 @@ func testClearSidebarOnTeamLeave(t *testing.T, ss store.Store, s SqlStore) { TeamID: team2.Id, ExcludeTeam: false, } - res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, err) require.NotEmpty(t, res) @@ -2334,6 +2349,8 @@ func testDeleteSidebarCategory(t *testing.T, ss store.Store, s SqlStore) { } func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) { + c := request.TestContext(t) + t.Run("Should be able to update sidebar channels", func(t *testing.T) { userId := model.NewId() teamId := model.NewId() @@ -2342,7 +2359,7 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) { TeamID: teamId, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) require.NoError(t, nErr) require.NotEmpty(t, res) @@ -2371,7 +2388,7 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) { TeamID: teamId, ExcludeTeam: false, } - res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts) + res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts) assert.NoError(t, nErr) require.NotEmpty(t, res) @@ -2391,6 +2408,8 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) { // in the hope of triggering a deadlock. This is a best-effort test case, and is not guaranteed // to catch a bug. func testSidebarCategoryDeadlock(t *testing.T, ss store.Store) { + c := request.TestContext(t) + userID := model.NewId() team := setupTeam(t, ss, userID) @@ -2413,7 +2432,7 @@ func testSidebarCategoryDeadlock(t *testing.T, ss store.Store) { TeamID: team.Id, ExcludeTeam: false, } - res, err := ss.Channel().CreateInitialSidebarCategories(userID, opts) + res, err := ss.Channel().CreateInitialSidebarCategories(c, userID, opts) require.NoError(t, err) require.NotEmpty(t, res) diff --git a/server/channels/store/storetest/compliance_store.go b/server/channels/store/storetest/compliance_store.go index d98adab045..ec581ca938 100644 --- a/server/channels/store/storetest/compliance_store.go +++ b/server/channels/store/storetest/compliance_store.go @@ -4,7 +4,6 @@ package storetest import ( - "context" "encoding/json" "testing" "time" @@ -13,6 +12,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -396,11 +396,13 @@ func testComplianceExportDirectMessages(t *testing.T, ss store.Store) { } func testMessageExportPublicChannel(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -470,7 +472,7 @@ func testMessageExportPublicChannel(t *testing.T, ss store.Store) { // fetch the message exports for both posts that user1 sent messageExportMap := map[string]model.MessageExport{} - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 2, len(messages)) @@ -500,11 +502,13 @@ func testMessageExportPublicChannel(t *testing.T, ss store.Store) { } func testMessageExportPrivateChannel(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -574,7 +578,7 @@ func testMessageExportPrivateChannel(t *testing.T, ss store.Store) { // fetch the message exports for both posts that user1 sent messageExportMap := map[string]model.MessageExport{} - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 2, len(messages)) @@ -606,11 +610,13 @@ func testMessageExportPrivateChannel(t *testing.T, ss store.Store) { } func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -665,7 +671,7 @@ func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) { // fetch the message export for the post that user1 sent messageExportMap := map[string]model.MessageExport{} - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) @@ -687,11 +693,13 @@ func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) { } func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -763,7 +771,7 @@ func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) { // fetch the message export for the post that user1 sent messageExportMap := map[string]model.MessageExport{} - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) @@ -785,10 +793,13 @@ func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) { // post,edit,export func testEditExportMessage(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) + // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -843,7 +854,7 @@ func testEditExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports from the start - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 2, len(messages)) @@ -877,10 +888,12 @@ func testEditExportMessage(t *testing.T, ss store.Store) { // post, export, edit, export func testEditAfterExportMessage(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -928,7 +941,7 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports from the start - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) @@ -954,7 +967,7 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports after edit - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: postEditTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: postEditTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 2, len(messages)) @@ -988,10 +1001,12 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) { // post, delete, export func testDeleteExportMessage(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -1044,7 +1059,7 @@ func testDeleteExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports from the start - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) @@ -1073,10 +1088,12 @@ func testDeleteExportMessage(t *testing.T, ss store.Store) { // post,export,delete,export func testDeleteAfterExportMessage(t *testing.T, ss store.Store) { + c := request.TestContext(t) + defer cleanupStoreState(t, ss) // get the starting number of message export entries startTime := model.GetMillis() - messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 0, len(messages)) @@ -1124,7 +1141,7 @@ func testDeleteAfterExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports from the start - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) @@ -1147,7 +1164,7 @@ func testDeleteAfterExportMessage(t *testing.T, ss store.Store) { require.NoError(t, err) // fetch the message exports after delete - messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: postDeleteTime - 1}, 10) + messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: postDeleteTime - 1}, 10) require.NoError(t, err) assert.Equal(t, 1, len(messages)) diff --git a/server/channels/store/storetest/license_store.go b/server/channels/store/storetest/license_store.go index e167ab8634..d991e4a2a5 100644 --- a/server/channels/store/storetest/license_store.go +++ b/server/channels/store/storetest/license_store.go @@ -4,12 +4,12 @@ package storetest import ( - "context" "testing" "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -36,6 +36,8 @@ func testLicenseStoreSave(t *testing.T, ss store.Store) { } func testLicenseStoreGet(t *testing.T, ss store.Store) { + c := request.TestContext(t) + l1 := model.LicenseRecord{} l1.Id = model.NewId() l1.Bytes = "junk" @@ -43,11 +45,11 @@ func testLicenseStoreGet(t *testing.T, ss store.Store) { err := ss.License().Save(&l1) require.NoError(t, err) - record, err := ss.License().Get(context.Background(), l1.Id) + record, err := ss.License().Get(c, l1.Id) require.NoError(t, err, "couldn't get license") require.Equal(t, record.Bytes, l1.Bytes, "license bytes didn't match") - _, err = ss.License().Get(context.Background(), "missing") + _, err = ss.License().Get(c, "missing") require.Error(t, err, "should fail on get license") } diff --git a/server/channels/store/storetest/mocks/ChannelStore.go b/server/channels/store/storetest/mocks/ChannelStore.go index 9bf9f33f92..cea58b35c2 100644 --- a/server/channels/store/storetest/mocks/ChannelStore.go +++ b/server/channels/store/storetest/mocks/ChannelStore.go @@ -10,6 +10,8 @@ import ( model "github.com/mattermost/mattermost/server/public/model" mock "github.com/stretchr/testify/mock" + request "github.com/mattermost/mattermost/server/public/shared/request" + store "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -270,25 +272,25 @@ func (_m *ChannelStore) CreateDirectChannel(userID *model.User, otherUserID *mod return r0, r1 } -// CreateInitialSidebarCategories provides a mock function with given fields: userID, opts -func (_m *ChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { - ret := _m.Called(userID, opts) +// CreateInitialSidebarCategories provides a mock function with given fields: c, userID, opts +func (_m *ChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { + ret := _m.Called(c, userID, opts) var r0 *model.OrderedSidebarCategories var r1 error - if rf, ok := ret.Get(0).(func(string, *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)); ok { - return rf(userID, opts) + if rf, ok := ret.Get(0).(func(request.CTX, string, *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)); ok { + return rf(c, userID, opts) } - if rf, ok := ret.Get(0).(func(string, *store.SidebarCategorySearchOpts) *model.OrderedSidebarCategories); ok { - r0 = rf(userID, opts) + if rf, ok := ret.Get(0).(func(request.CTX, string, *store.SidebarCategorySearchOpts) *model.OrderedSidebarCategories); ok { + r0 = rf(c, userID, opts) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.OrderedSidebarCategories) } } - if rf, ok := ret.Get(1).(func(string, *store.SidebarCategorySearchOpts) error); ok { - r1 = rf(userID, opts) + if rf, ok := ret.Get(1).(func(request.CTX, string, *store.SidebarCategorySearchOpts) error); ok { + r1 = rf(c, userID, opts) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/ComplianceStore.go b/server/channels/store/storetest/mocks/ComplianceStore.go index 2e16d49024..ecd1a0e070 100644 --- a/server/channels/store/storetest/mocks/ComplianceStore.go +++ b/server/channels/store/storetest/mocks/ComplianceStore.go @@ -5,9 +5,8 @@ package mocks import ( - context "context" - model "github.com/mattermost/mattermost/server/public/model" + request "github.com/mattermost/mattermost/server/public/shared/request" mock "github.com/stretchr/testify/mock" ) @@ -101,32 +100,32 @@ func (_m *ComplianceStore) GetAll(offset int, limit int) (model.Compliances, err return r0, r1 } -// MessageExport provides a mock function with given fields: ctx, cursor, limit -func (_m *ComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { - ret := _m.Called(ctx, cursor, limit) +// MessageExport provides a mock function with given fields: c, cursor, limit +func (_m *ComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { + ret := _m.Called(c, cursor, limit) var r0 []*model.MessageExport var r1 model.MessageExportCursor var r2 error - if rf, ok := ret.Get(0).(func(context.Context, model.MessageExportCursor, int) ([]*model.MessageExport, model.MessageExportCursor, error)); ok { - return rf(ctx, cursor, limit) + if rf, ok := ret.Get(0).(func(request.CTX, model.MessageExportCursor, int) ([]*model.MessageExport, model.MessageExportCursor, error)); ok { + return rf(c, cursor, limit) } - if rf, ok := ret.Get(0).(func(context.Context, model.MessageExportCursor, int) []*model.MessageExport); ok { - r0 = rf(ctx, cursor, limit) + if rf, ok := ret.Get(0).(func(request.CTX, model.MessageExportCursor, int) []*model.MessageExport); ok { + r0 = rf(c, cursor, limit) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*model.MessageExport) } } - if rf, ok := ret.Get(1).(func(context.Context, model.MessageExportCursor, int) model.MessageExportCursor); ok { - r1 = rf(ctx, cursor, limit) + if rf, ok := ret.Get(1).(func(request.CTX, model.MessageExportCursor, int) model.MessageExportCursor); ok { + r1 = rf(c, cursor, limit) } else { r1 = ret.Get(1).(model.MessageExportCursor) } - if rf, ok := ret.Get(2).(func(context.Context, model.MessageExportCursor, int) error); ok { - r2 = rf(ctx, cursor, limit) + if rf, ok := ret.Get(2).(func(request.CTX, model.MessageExportCursor, int) error); ok { + r2 = rf(c, cursor, limit) } else { r2 = ret.Error(2) } diff --git a/server/channels/store/storetest/mocks/EmojiStore.go b/server/channels/store/storetest/mocks/EmojiStore.go index 58e25fd838..1e830a15a4 100644 --- a/server/channels/store/storetest/mocks/EmojiStore.go +++ b/server/channels/store/storetest/mocks/EmojiStore.go @@ -29,17 +29,17 @@ func (_m *EmojiStore) Delete(emoji *model.Emoji, timestamp int64) error { return r0 } -// Get provides a mock function with given fields: ctx, id, allowFromCache -func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { - ret := _m.Called(ctx, id, allowFromCache) +// Get provides a mock function with given fields: c, id, allowFromCache +func (_m *EmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { + ret := _m.Called(c, id, allowFromCache) var r0 *model.Emoji var r1 error if rf, ok := ret.Get(0).(func(request.CTX, string, bool) (*model.Emoji, error)); ok { - return rf(ctx, id, allowFromCache) + return rf(c, id, allowFromCache) } if rf, ok := ret.Get(0).(func(request.CTX, string, bool) *model.Emoji); ok { - r0 = rf(ctx, id, allowFromCache) + r0 = rf(c, id, allowFromCache) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Emoji) @@ -47,7 +47,7 @@ func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*mod } if rf, ok := ret.Get(1).(func(request.CTX, string, bool) error); ok { - r1 = rf(ctx, id, allowFromCache) + r1 = rf(c, id, allowFromCache) } else { r1 = ret.Error(1) } @@ -55,17 +55,17 @@ func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*mod return r0, r1 } -// GetByName provides a mock function with given fields: ctx, name, allowFromCache -func (_m *EmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { - ret := _m.Called(ctx, name, allowFromCache) +// GetByName provides a mock function with given fields: c, name, allowFromCache +func (_m *EmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { + ret := _m.Called(c, name, allowFromCache) var r0 *model.Emoji var r1 error if rf, ok := ret.Get(0).(func(request.CTX, string, bool) (*model.Emoji, error)); ok { - return rf(ctx, name, allowFromCache) + return rf(c, name, allowFromCache) } if rf, ok := ret.Get(0).(func(request.CTX, string, bool) *model.Emoji); ok { - r0 = rf(ctx, name, allowFromCache) + r0 = rf(c, name, allowFromCache) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Emoji) @@ -73,7 +73,7 @@ func (_m *EmojiStore) GetByName(ctx request.CTX, name string, allowFromCache boo } if rf, ok := ret.Get(1).(func(request.CTX, string, bool) error); ok { - r1 = rf(ctx, name, allowFromCache) + r1 = rf(c, name, allowFromCache) } else { r1 = ret.Error(1) } @@ -107,17 +107,17 @@ func (_m *EmojiStore) GetList(offset int, limit int, sort string) ([]*model.Emoj return r0, r1 } -// GetMultipleByName provides a mock function with given fields: ctx, names -func (_m *EmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) { - ret := _m.Called(ctx, names) +// GetMultipleByName provides a mock function with given fields: c, names +func (_m *EmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) { + ret := _m.Called(c, names) var r0 []*model.Emoji var r1 error if rf, ok := ret.Get(0).(func(request.CTX, []string) ([]*model.Emoji, error)); ok { - return rf(ctx, names) + return rf(c, names) } if rf, ok := ret.Get(0).(func(request.CTX, []string) []*model.Emoji); ok { - r0 = rf(ctx, names) + r0 = rf(c, names) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*model.Emoji) @@ -125,7 +125,7 @@ func (_m *EmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*mod } if rf, ok := ret.Get(1).(func(request.CTX, []string) error); ok { - r1 = rf(ctx, names) + r1 = rf(c, names) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/LicenseStore.go b/server/channels/store/storetest/mocks/LicenseStore.go index 04c9ddd764..2f765dbaf4 100644 --- a/server/channels/store/storetest/mocks/LicenseStore.go +++ b/server/channels/store/storetest/mocks/LicenseStore.go @@ -5,9 +5,8 @@ package mocks import ( - context "context" - model "github.com/mattermost/mattermost/server/public/model" + request "github.com/mattermost/mattermost/server/public/shared/request" mock "github.com/stretchr/testify/mock" ) @@ -16,25 +15,25 @@ type LicenseStore struct { mock.Mock } -// Get provides a mock function with given fields: ctx, id -func (_m *LicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) { - ret := _m.Called(ctx, id) +// Get provides a mock function with given fields: c, id +func (_m *LicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) { + ret := _m.Called(c, id) var r0 *model.LicenseRecord var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string) (*model.LicenseRecord, error)); ok { - return rf(ctx, id) + if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.LicenseRecord, error)); ok { + return rf(c, id) } - if rf, ok := ret.Get(0).(func(context.Context, string) *model.LicenseRecord); ok { - r0 = rf(ctx, id) + if rf, ok := ret.Get(0).(func(request.CTX, string) *model.LicenseRecord); ok { + r0 = rf(c, id) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.LicenseRecord) } } - if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = rf(ctx, id) + if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok { + r1 = rf(c, id) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/SessionStore.go b/server/channels/store/storetest/mocks/SessionStore.go index fe2fe13715..43928f0488 100644 --- a/server/channels/store/storetest/mocks/SessionStore.go +++ b/server/channels/store/storetest/mocks/SessionStore.go @@ -5,9 +5,8 @@ package mocks import ( - context "context" - model "github.com/mattermost/mattermost/server/public/model" + request "github.com/mattermost/mattermost/server/public/shared/request" mock "github.com/stretchr/testify/mock" ) @@ -54,25 +53,25 @@ func (_m *SessionStore) Cleanup(expiryTime int64, batchSize int64) error { return r0 } -// Get provides a mock function with given fields: ctx, sessionIDOrToken -func (_m *SessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) { - ret := _m.Called(ctx, sessionIDOrToken) +// Get provides a mock function with given fields: c, sessionIDOrToken +func (_m *SessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) { + ret := _m.Called(c, sessionIDOrToken) var r0 *model.Session var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string) (*model.Session, error)); ok { - return rf(ctx, sessionIDOrToken) + if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.Session, error)); ok { + return rf(c, sessionIDOrToken) } - if rf, ok := ret.Get(0).(func(context.Context, string) *model.Session); ok { - r0 = rf(ctx, sessionIDOrToken) + if rf, ok := ret.Get(0).(func(request.CTX, string) *model.Session); ok { + r0 = rf(c, sessionIDOrToken) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Session) } } - if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = rf(ctx, sessionIDOrToken) + if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok { + r1 = rf(c, sessionIDOrToken) } else { r1 = ret.Error(1) } @@ -80,25 +79,25 @@ func (_m *SessionStore) Get(ctx context.Context, sessionIDOrToken string) (*mode return r0, r1 } -// GetSessions provides a mock function with given fields: userID -func (_m *SessionStore) GetSessions(userID string) ([]*model.Session, error) { - ret := _m.Called(userID) +// GetSessions provides a mock function with given fields: c, userID +func (_m *SessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) { + ret := _m.Called(c, userID) var r0 []*model.Session var r1 error - if rf, ok := ret.Get(0).(func(string) ([]*model.Session, error)); ok { - return rf(userID) + if rf, ok := ret.Get(0).(func(*request.Context, string) ([]*model.Session, error)); ok { + return rf(c, userID) } - if rf, ok := ret.Get(0).(func(string) []*model.Session); ok { - r0 = rf(userID) + if rf, ok := ret.Get(0).(func(*request.Context, string) []*model.Session); ok { + r0 = rf(c, userID) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*model.Session) } } - if rf, ok := ret.Get(1).(func(string) error); ok { - r1 = rf(userID) + if rf, ok := ret.Get(1).(func(*request.Context, string) error); ok { + r1 = rf(c, userID) } else { r1 = ret.Error(1) } @@ -200,25 +199,25 @@ func (_m *SessionStore) RemoveAllSessions() error { return r0 } -// Save provides a mock function with given fields: session -func (_m *SessionStore) Save(session *model.Session) (*model.Session, error) { - ret := _m.Called(session) +// Save provides a mock function with given fields: c, session +func (_m *SessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) { + ret := _m.Called(c, session) var r0 *model.Session var r1 error - if rf, ok := ret.Get(0).(func(*model.Session) (*model.Session, error)); ok { - return rf(session) + if rf, ok := ret.Get(0).(func(request.CTX, *model.Session) (*model.Session, error)); ok { + return rf(c, session) } - if rf, ok := ret.Get(0).(func(*model.Session) *model.Session); ok { - r0 = rf(session) + if rf, ok := ret.Get(0).(func(request.CTX, *model.Session) *model.Session); ok { + r0 = rf(c, session) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Session) } } - if rf, ok := ret.Get(1).(func(*model.Session) error); ok { - r1 = rf(session) + if rf, ok := ret.Get(1).(func(request.CTX, *model.Session) error); ok { + r1 = rf(c, session) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/TeamStore.go b/server/channels/store/storetest/mocks/TeamStore.go index 110886c2f9..bcc64f1ee2 100644 --- a/server/channels/store/storetest/mocks/TeamStore.go +++ b/server/channels/store/storetest/mocks/TeamStore.go @@ -5,9 +5,8 @@ package mocks import ( - context "context" - model "github.com/mattermost/mattermost/server/public/model" + request "github.com/mattermost/mattermost/server/public/shared/request" mock "github.com/stretchr/testify/mock" ) @@ -497,25 +496,25 @@ func (_m *TeamStore) GetMany(ids []string) ([]*model.Team, error) { 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) +// GetMember provides a mock function with given fields: c, teamID, userID +func (_m *TeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) { + ret := _m.Called(c, teamID, userID) var r0 *model.TeamMember var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string, string) (*model.TeamMember, error)); ok { - return rf(ctx, teamID, userID) + if rf, ok := ret.Get(0).(func(request.CTX, string, string) (*model.TeamMember, error)); ok { + return rf(c, teamID, userID) } - if rf, ok := ret.Get(0).(func(context.Context, string, string) *model.TeamMember); ok { - r0 = rf(ctx, teamID, userID) + if rf, ok := ret.Get(0).(func(request.CTX, string, string) *model.TeamMember); ok { + r0 = rf(c, teamID, userID) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.TeamMember) } } - if rf, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = rf(ctx, teamID, userID) + if rf, ok := ret.Get(1).(func(request.CTX, string, string) error); ok { + r1 = rf(c, teamID, userID) } else { r1 = ret.Error(1) } @@ -653,25 +652,25 @@ func (_m *TeamStore) GetTeamsByUserId(userID string) ([]*model.Team, error) { return r0, r1 } -// GetTeamsForUser provides a mock function with given fields: ctx, userID, excludeTeamID, includeDeleted -func (_m *TeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { - ret := _m.Called(ctx, userID, excludeTeamID, includeDeleted) +// GetTeamsForUser provides a mock function with given fields: c, userID, excludeTeamID, includeDeleted +func (_m *TeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { + ret := _m.Called(c, userID, excludeTeamID, includeDeleted) var r0 []*model.TeamMember var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string, string, bool) ([]*model.TeamMember, error)); ok { - return rf(ctx, userID, excludeTeamID, includeDeleted) + if rf, ok := ret.Get(0).(func(request.CTX, string, string, bool) ([]*model.TeamMember, error)); ok { + return rf(c, userID, excludeTeamID, includeDeleted) } - if rf, ok := ret.Get(0).(func(context.Context, string, string, bool) []*model.TeamMember); ok { - r0 = rf(ctx, userID, excludeTeamID, includeDeleted) + if rf, ok := ret.Get(0).(func(request.CTX, string, string, bool) []*model.TeamMember); ok { + r0 = rf(c, userID, excludeTeamID, includeDeleted) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*model.TeamMember) } } - if rf, ok := ret.Get(1).(func(context.Context, string, string, bool) error); ok { - r1 = rf(ctx, userID, excludeTeamID, includeDeleted) + if rf, ok := ret.Get(1).(func(request.CTX, string, string, bool) error); ok { + r1 = rf(c, userID, excludeTeamID, includeDeleted) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/mocks/UploadSessionStore.go b/server/channels/store/storetest/mocks/UploadSessionStore.go index 267d50c160..05fcef13e9 100644 --- a/server/channels/store/storetest/mocks/UploadSessionStore.go +++ b/server/channels/store/storetest/mocks/UploadSessionStore.go @@ -5,9 +5,8 @@ package mocks import ( - context "context" - model "github.com/mattermost/mattermost/server/public/model" + request "github.com/mattermost/mattermost/server/public/shared/request" mock "github.com/stretchr/testify/mock" ) @@ -30,25 +29,25 @@ func (_m *UploadSessionStore) Delete(id string) error { return r0 } -// Get provides a mock function with given fields: ctx, id -func (_m *UploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { - ret := _m.Called(ctx, id) +// Get provides a mock function with given fields: c, id +func (_m *UploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) { + ret := _m.Called(c, id) var r0 *model.UploadSession var r1 error - if rf, ok := ret.Get(0).(func(context.Context, string) (*model.UploadSession, error)); ok { - return rf(ctx, id) + if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.UploadSession, error)); ok { + return rf(c, id) } - if rf, ok := ret.Get(0).(func(context.Context, string) *model.UploadSession); ok { - r0 = rf(ctx, id) + if rf, ok := ret.Get(0).(func(request.CTX, string) *model.UploadSession); ok { + r0 = rf(c, id) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.UploadSession) } } - if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = rf(ctx, id) + if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok { + r1 = rf(c, id) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/oauth_store.go b/server/channels/store/storetest/oauth_store.go index d58fb84af0..6dd12ba13c 100644 --- a/server/channels/store/storetest/oauth_store.go +++ b/server/channels/store/storetest/oauth_store.go @@ -4,13 +4,13 @@ package storetest import ( - "context" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -357,6 +357,8 @@ func testOAuthGetAccessDataByUserForApp(t *testing.T, ss store.Store) { } func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) { + c := request.TestContext(t) + a1 := model.OAuthApp{} a1.CreatorId = model.NewId() a1.Name = "TestApp" + model.NewId() @@ -374,7 +376,7 @@ func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) { s1.Token = model.NewId() s1.IsOAuth = true - s1, nErr := ss.Session().Save(s1) + s1, nErr := ss.Session().Save(c, s1) require.NoError(t, nErr) ad1 := model.AccessData{} @@ -390,7 +392,7 @@ func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) { err = ss.OAuth().DeleteApp(a1.Id) require.NoError(t, err) - _, nErr = ss.Session().Get(context.Background(), s1.Token) + _, nErr = ss.Session().Get(c, s1.Token) require.Error(t, nErr, "should error - session should be deleted") _, err = ss.OAuth().GetAccessData(s1.Token) diff --git a/server/channels/store/storetest/session_store.go b/server/channels/store/storetest/session_store.go index e6e86a33c6..5503a6d58a 100644 --- a/server/channels/store/storetest/session_store.go +++ b/server/channels/store/storetest/session_store.go @@ -4,13 +4,13 @@ package storetest import ( - "context" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -39,34 +39,38 @@ func TestSessionStore(t *testing.T, ss store.Store) { } func testSessionStoreSave(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - _, err := ss.Session().Save(s1) + _, err := ss.Session().Save(c, s1) require.NoError(t, err) } func testSessionGet(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) s2 := &model.Session{} s2.UserId = s1.UserId - _, err = ss.Session().Save(s2) + _, err = ss.Session().Save(c, s2) require.NoError(t, err) s3 := &model.Session{} s3.UserId = s1.UserId s3.ExpiresAt = 1 - _, err = ss.Session().Save(s3) + _, err = ss.Session().Save(c, s3) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.Equal(t, session.Id, s1.Id, "should match") @@ -75,21 +79,23 @@ func testSessionGet(t *testing.T, ss store.Store) { err = ss.Session().UpdateProps(session) require.NoError(t, err) - session2, err := ss.Session().Get(context.Background(), session.Id) + session2, err := ss.Session().Get(c, session.Id) require.NoError(t, err) require.Equal(t, session.Props, session2.Props, "should match") - data, err := ss.Session().GetSessions(s1.UserId) + data, err := ss.Session().GetSessions(c, s1.UserId) require.NoError(t, err) require.Len(t, data, 3, "should match len") } func testSessionGetWithDeviceId(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() s1.ExpiresAt = model.GetMillis() + 10000 - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) s2 := &model.Session{} @@ -97,7 +103,7 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) { s2.DeviceId = model.NewId() s2.ExpiresAt = model.GetMillis() + 10000 - _, err = ss.Session().Save(s2) + _, err = ss.Session().Save(c, s2) require.NoError(t, err) s3 := &model.Session{} @@ -105,7 +111,7 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) { s3.ExpiresAt = 1 s3.DeviceId = model.NewId() - _, err = ss.Session().Save(s3) + _, err = ss.Session().Save(c, s3) require.NoError(t, err) data, err := ss.Session().GetSessionsWithActiveDeviceIds(s1.UserId) @@ -114,86 +120,96 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) { } func testSessionRemove(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.Equal(t, session.Id, s1.Id, "should match") removeErr := ss.Session().Remove(s1.Id) require.NoError(t, removeErr) - _, err = ss.Session().Get(context.Background(), s1.Id) + _, err = ss.Session().Get(c, s1.Id) require.Error(t, err, "should have been removed") } func testSessionRemoveAll(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.Equal(t, session.Id, s1.Id, "should match") removeErr := ss.Session().RemoveAllSessions() require.NoError(t, removeErr) - _, err = ss.Session().Get(context.Background(), s1.Id) + _, err = ss.Session().Get(c, s1.Id) require.Error(t, err, "should have been removed") } func testSessionRemoveByUser(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.Equal(t, session.Id, s1.Id, "should match") deleteErr := ss.Session().PermanentDeleteSessionsByUser(s1.UserId) require.NoError(t, deleteErr) - _, err = ss.Session().Get(context.Background(), s1.Id) + _, err = ss.Session().Get(c, s1.Id) require.Error(t, err, "should have been removed") } func testSessionRemoveToken(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.Equal(t, session.Id, s1.Id, "should match") removeErr := ss.Session().Remove(s1.Token) require.NoError(t, removeErr) - _, err = ss.Session().Get(context.Background(), s1.Id) + _, err = ss.Session().Get(c, s1.Id) require.Error(t, err, "should have been removed") - data, err := ss.Session().GetSessions(s1.UserId) + data, err := ss.Session().GetSessions(c, s1.UserId) require.NoError(t, err) require.Empty(t, data, "should match len") } func testSessionUpdateDeviceId(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) _, err = ss.Session().UpdateDeviceId(s1.Id, model.PushNotifyApple+":1234567890", s1.ExpiresAt) @@ -202,7 +218,7 @@ func testSessionUpdateDeviceId(t *testing.T, ss store.Store) { s2 := &model.Session{} s2.UserId = model.NewId() - s2, err = ss.Session().Save(s2) + s2, err = ss.Session().Save(c, s2) require.NoError(t, err) _, err = ss.Session().UpdateDeviceId(s2.Id, model.PushNotifyApple+":1234567890", s1.ExpiresAt) @@ -210,10 +226,12 @@ func testSessionUpdateDeviceId(t *testing.T, ss store.Store) { } func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) _, err = ss.Session().UpdateDeviceId(s1.Id, model.PushNotifyAppleReactNative+":1234567890", s1.ExpiresAt) @@ -222,7 +240,7 @@ func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) { s2 := &model.Session{} s2.UserId = model.NewId() - s2, err = ss.Session().Save(s2) + s2, err = ss.Session().Save(c, s2) require.NoError(t, err) _, err = ss.Session().UpdateDeviceId(s2.Id, model.PushNotifyAppleReactNative+":1234567890", s1.ExpiresAt) @@ -230,41 +248,47 @@ func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) { } func testSessionStoreUpdateExpiresAt(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) err = ss.Session().UpdateExpiresAt(s1.Id, 1234567890) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.EqualValues(t, session.ExpiresAt, 1234567890, "ExpiresAt not updated correctly") } func testSessionStoreUpdateLastActivityAt(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) err = ss.Session().UpdateLastActivityAt(s1.Id, 1234567890) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.EqualValues(t, session.LastActivityAt, 1234567890, "LastActivityAt not updated correctly") } func testSessionCount(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() s1.ExpiresAt = model.GetMillis() + 100000 - _, err := ss.Session().Save(s1) + _, err := ss.Session().Save(c, s1) require.NoError(t, err) count, err := ss.Session().AnalyticsSessionCount() @@ -273,49 +297,51 @@ func testSessionCount(t *testing.T, ss store.Store) { } func testSessionCleanup(t *testing.T, ss store.Store) { + c := request.TestContext(t) + now := model.GetMillis() s1 := &model.Session{} s1.UserId = model.NewId() s1.ExpiresAt = 0 // never expires - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) s2 := &model.Session{} s2.UserId = s1.UserId s2.ExpiresAt = now + 1000000 // expires in the future - s2, err = ss.Session().Save(s2) + s2, err = ss.Session().Save(c, s2) require.NoError(t, err) s3 := &model.Session{} s3.UserId = model.NewId() s3.ExpiresAt = 1 // expired - s3, err = ss.Session().Save(s3) + s3, err = ss.Session().Save(c, s3) require.NoError(t, err) s4 := &model.Session{} s4.UserId = model.NewId() s4.ExpiresAt = 2 // expired - s4, err = ss.Session().Save(s4) + s4, err = ss.Session().Save(c, s4) require.NoError(t, err) err = ss.Session().Cleanup(now, 1) require.NoError(t, err) - _, err = ss.Session().Get(context.Background(), s1.Id) + _, err = ss.Session().Get(c, s1.Id) assert.NoError(t, err) - _, err = ss.Session().Get(context.Background(), s2.Id) + _, err = ss.Session().Get(c, s2.Id) assert.NoError(t, err) - _, err = ss.Session().Get(context.Background(), s3.Id) + _, err = ss.Session().Get(c, s3.Id) assert.Error(t, err) - _, err = ss.Session().Get(context.Background(), s4.Id) + _, err = ss.Session().Get(c, s4.Id) assert.Error(t, err) removeErr := ss.Session().Remove(s1.Id) @@ -326,6 +352,8 @@ func testSessionCleanup(t *testing.T, ss store.Store) { } func testGetSessionsExpired(t *testing.T, ss store.Store) { + c := request.TestContext(t) + now := model.GetMillis() // Clear existing sessions. @@ -336,34 +364,34 @@ func testGetSessionsExpired(t *testing.T, ss store.Store) { s1.UserId = model.NewId() s1.DeviceId = model.NewId() s1.ExpiresAt = 0 // never expires - _, err = ss.Session().Save(s1) + _, err = ss.Session().Save(c, s1) require.NoError(t, err) s2 := &model.Session{} s2.UserId = model.NewId() s2.DeviceId = model.NewId() s2.ExpiresAt = now - TenMinutes // expired within threshold - s2, err = ss.Session().Save(s2) + s2, err = ss.Session().Save(c, s2) require.NoError(t, err) s3 := &model.Session{} s3.UserId = model.NewId() s3.DeviceId = model.NewId() s3.ExpiresAt = now - (TenMinutes * 100) // expired outside threshold - _, err = ss.Session().Save(s3) + _, err = ss.Session().Save(c, s3) require.NoError(t, err) s4 := &model.Session{} s4.UserId = model.NewId() s4.ExpiresAt = now - TenMinutes // expired within threshold, but not mobile - s4, err = ss.Session().Save(s4) + s4, err = ss.Session().Save(c, s4) require.NoError(t, err) s5 := &model.Session{} s5.UserId = model.NewId() s5.DeviceId = model.NewId() s5.ExpiresAt = now + (TenMinutes * 100000) // not expired - _, err = ss.Session().Save(s5) + _, err = ss.Session().Save(c, s5) require.NoError(t, err) sessions, err := ss.Session().GetSessionsExpired(TenMinutes*2, true, true) // mobile only @@ -381,26 +409,28 @@ func testGetSessionsExpired(t *testing.T, ss store.Store) { } func testUpdateExpiredNotify(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := &model.Session{} s1.UserId = model.NewId() s1.DeviceId = model.NewId() s1.ExpiresAt = model.GetMillis() + TenMinutes - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) - session, err := ss.Session().Get(context.Background(), s1.Id) + session, err := ss.Session().Get(c, s1.Id) require.NoError(t, err) require.False(t, session.ExpiredNotify) err = ss.Session().UpdateExpiredNotify(session.Id, true) require.NoError(t, err) - session, err = ss.Session().Get(context.Background(), s1.Id) + session, err = ss.Session().Get(c, s1.Id) require.NoError(t, err) require.True(t, session.ExpiredNotify) err = ss.Session().UpdateExpiredNotify(session.Id, false) require.NoError(t, err) - session, err = ss.Session().Get(context.Background(), s1.Id) + session, err = ss.Session().Get(c, s1.Id) require.NoError(t, err) require.False(t, session.ExpiredNotify) } diff --git a/server/channels/store/storetest/team_store.go b/server/channels/store/storetest/team_store.go index cc54bb7f9f..08e089e218 100644 --- a/server/channels/store/storetest/team_store.go +++ b/server/channels/store/storetest/team_store.go @@ -14,6 +14,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -1299,6 +1300,8 @@ func testGetMembers(t *testing.T, ss store.Store) { } func testTeamMembers(t *testing.T, ss store.Store) { + c := request.TestContext(t) + teamId1 := model.NewId() teamId2 := model.NewId() @@ -1318,8 +1321,7 @@ func testTeamMembers(t *testing.T, ss store.Store) { require.Len(t, ms, 1) require.Equal(t, m3.UserId, ms[0].UserId) - ctx := context.Background() - ms, err = ss.Team().GetTeamsForUser(ctx, m1.UserId, "", true) + ms, err = ss.Team().GetTeamsForUser(c, m1.UserId, "", true) require.NoError(t, err) require.Len(t, ms, 1) require.Equal(t, m1.TeamId, ms[0].TeamId) @@ -1348,11 +1350,11 @@ func testTeamMembers(t *testing.T, ss store.Store) { _, nErr = ss.Team().SaveMultipleMembers([]*model.TeamMember{m4, m5}, -1) require.NoError(t, nErr) - ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", true) + ms, err = ss.Team().GetTeamsForUser(c, uid, "", true) require.NoError(t, err) require.Len(t, ms, 2) - ms, err = ss.Team().GetTeamsForUser(ctx, uid, teamId2, true) + ms, err = ss.Team().GetTeamsForUser(c, uid, teamId2, true) require.NoError(t, err) require.Len(t, ms, 1) @@ -1360,18 +1362,18 @@ func testTeamMembers(t *testing.T, ss store.Store) { _, err = ss.Team().UpdateMember(m4) require.NoError(t, err) - ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", true) + ms, err = ss.Team().GetTeamsForUser(c, uid, "", true) require.NoError(t, err) require.Len(t, ms, 2) - ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", false) + ms, err = ss.Team().GetTeamsForUser(c, uid, "", false) require.NoError(t, err) require.Len(t, ms, 1) nErr = ss.Team().RemoveAllMembersByUser(uid) require.NoError(t, nErr) - ms, err = ss.Team().GetTeamsForUser(ctx, m1.UserId, "", true) + ms, err = ss.Team().GetTeamsForUser(c, m1.UserId, "", true) require.NoError(t, err) require.Empty(t, ms) } @@ -2974,6 +2976,8 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { } func testGetTeamMember(t *testing.T, ss store.Store) { + c := request.TestContext(t) + teamId1 := model.NewId() m1 := &model.TeamMember{TeamId: teamId1, UserId: model.NewId()} @@ -2981,17 +2985,17 @@ func testGetTeamMember(t *testing.T, ss store.Store) { require.NoError(t, nErr) var rm1 *model.TeamMember - rm1, err := ss.Team().GetMember(context.Background(), m1.TeamId, m1.UserId) + rm1, err := ss.Team().GetMember(c, m1.TeamId, m1.UserId) require.NoError(t, err) require.Equal(t, rm1.TeamId, m1.TeamId, "bad team id") require.Equal(t, rm1.UserId, m1.UserId, "bad user id") - _, err = ss.Team().GetMember(context.Background(), m1.TeamId, "") + _, err = ss.Team().GetMember(c, m1.TeamId, "") require.Error(t, err, "empty user id - should have failed") - _, err = ss.Team().GetMember(context.Background(), "", m1.UserId) + _, err = ss.Team().GetMember(c, "", m1.UserId) require.Error(t, err, "empty team id - should have failed") // Test with a custom team scheme. @@ -3021,7 +3025,7 @@ func testGetTeamMember(t *testing.T, ss store.Store) { _, nErr = ss.Team().SaveMember(m2, -1) require.NoError(t, nErr) - m3, err := ss.Team().GetMember(context.Background(), m2.TeamId, m2.UserId) + m3, err := ss.Team().GetMember(c, m2.TeamId, m2.UserId) require.NoError(t, err) t.Log(m3) @@ -3031,7 +3035,7 @@ func testGetTeamMember(t *testing.T, ss store.Store) { _, nErr = ss.Team().SaveMember(m4, -1) require.NoError(t, nErr) - m5, err := ss.Team().GetMember(context.Background(), m4.TeamId, m4.UserId) + m5, err := ss.Team().GetMember(c, m4.TeamId, m4.UserId) require.NoError(t, err) assert.Equal(t, s2.DefaultTeamGuestRole, m5.Roles) @@ -3292,6 +3296,8 @@ func testGetTeamsByScheme(t *testing.T, ss store.Store) { } func testTeamStoreMigrateTeamMembers(t *testing.T, ss store.Store) { + c := request.TestContext(t) + s1 := model.NewId() t1 := &model.Team{ DisplayName: "Name", @@ -3341,19 +3347,19 @@ func testTeamStoreMigrateTeamMembers(t *testing.T, ss store.Store) { } } - tm1b, err := ss.Team().GetMember(context.Background(), tm1.TeamId, tm1.UserId) + tm1b, err := ss.Team().GetMember(c, tm1.TeamId, tm1.UserId) assert.NoError(t, err) assert.Equal(t, "", tm1b.ExplicitRoles) assert.True(t, tm1b.SchemeUser) assert.True(t, tm1b.SchemeAdmin) - tm2b, err := ss.Team().GetMember(context.Background(), tm2.TeamId, tm2.UserId) + tm2b, err := ss.Team().GetMember(c, tm2.TeamId, tm2.UserId) assert.NoError(t, err) assert.Equal(t, "", tm2b.ExplicitRoles) assert.True(t, tm2b.SchemeUser) assert.False(t, tm2b.SchemeAdmin) - tm3b, err := ss.Team().GetMember(context.Background(), tm3.TeamId, tm3.UserId) + tm3b, err := ss.Team().GetMember(c, tm3.TeamId, tm3.UserId) assert.NoError(t, err) assert.Equal(t, "something_else", tm3b.ExplicitRoles) assert.False(t, tm3b.SchemeUser) @@ -3408,6 +3414,8 @@ func testResetAllTeamSchemes(t *testing.T, ss store.Store) { } func testTeamStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) { + c := request.TestContext(t) + m1 := &model.TeamMember{ TeamId: model.NewId(), UserId: model.NewId(), @@ -3434,19 +3442,19 @@ func testTeamStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) { require.NoError(t, (ss.Team().ClearAllCustomRoleAssignments())) - r1, err := ss.Team().GetMember(context.Background(), m1.TeamId, m1.UserId) + r1, err := ss.Team().GetMember(c, m1.TeamId, m1.UserId) require.NoError(t, err) assert.Equal(t, m1.ExplicitRoles, r1.Roles) - r2, err := ss.Team().GetMember(context.Background(), m2.TeamId, m2.UserId) + r2, err := ss.Team().GetMember(c, m2.TeamId, m2.UserId) require.NoError(t, err) assert.Equal(t, "team_user team_admin", r2.Roles) - r3, err := ss.Team().GetMember(context.Background(), m3.TeamId, m3.UserId) + r3, err := ss.Team().GetMember(c, m3.TeamId, m3.UserId) require.NoError(t, err) assert.Equal(t, m3.ExplicitRoles, r3.Roles) - r4, err := ss.Team().GetMember(context.Background(), m4.TeamId, m4.UserId) + r4, err := ss.Team().GetMember(c, m4.TeamId, m4.UserId) require.NoError(t, err) assert.Equal(t, "", r4.Roles) } diff --git a/server/channels/store/storetest/upload_session_store.go b/server/channels/store/storetest/upload_session_store.go index 21e11d46c2..accb0b0e38 100644 --- a/server/channels/store/storetest/upload_session_store.go +++ b/server/channels/store/storetest/upload_session_store.go @@ -4,13 +4,13 @@ package storetest import ( - "context" "testing" "time" "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -22,6 +22,8 @@ func TestUploadSessionStore(t *testing.T, ss store.Store) { } func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) { + c := request.TestContext(t) + var session *model.UploadSession t.Run("saving nil session should fail", func(t *testing.T) { @@ -53,13 +55,13 @@ func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) { }) t.Run("getting non-existing session should fail", func(t *testing.T) { - us, err := ss.UploadSession().Get(context.Background(), "fake") + us, err := ss.UploadSession().Get(c, "fake") require.Error(t, err) require.Nil(t, us) }) t.Run("getting existing session should succeed", func(t *testing.T) { - us, err := ss.UploadSession().Get(context.Background(), session.Id) + us, err := ss.UploadSession().Get(c, session.Id) require.NoError(t, err) require.NotNil(t, us) require.Equal(t, session, us) @@ -67,6 +69,8 @@ func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) { } func testUploadSessionStoreUpdate(t *testing.T, ss store.Store) { + c := request.TestContext(t) + session := &model.UploadSession{ Type: model.UploadTypeAttachment, UserId: model.NewId(), @@ -101,7 +105,7 @@ func testUploadSessionStoreUpdate(t *testing.T, ss store.Store) { err = ss.UploadSession().Update(us) require.NoError(t, err) - updated, err := ss.UploadSession().Get(context.Background(), us.Id) + updated, err := ss.UploadSession().Get(c, us.Id) require.NoError(t, err) require.NotNil(t, us) require.Equal(t, us, updated) @@ -176,6 +180,8 @@ func testUploadSessionStoreGetForUser(t *testing.T, ss store.Store) { } func testUploadSessionStoreDelete(t *testing.T, ss store.Store) { + c := request.TestContext(t) + session := &model.UploadSession{ Id: model.NewId(), Type: model.UploadTypeAttachment, @@ -200,7 +206,7 @@ func testUploadSessionStoreDelete(t *testing.T, ss store.Store) { err = ss.UploadSession().Delete(session.Id) require.NoError(t, err) - us, err = ss.UploadSession().Get(context.Background(), us.Id) + us, err = ss.UploadSession().Get(c, us.Id) require.Error(t, err) require.Nil(t, us) require.IsType(t, &store.ErrNotFound{}, err) diff --git a/server/channels/store/storetest/user_access_token_store.go b/server/channels/store/storetest/user_access_token_store.go index fad63163d7..a3aed42617 100644 --- a/server/channels/store/storetest/user_access_token_store.go +++ b/server/channels/store/storetest/user_access_token_store.go @@ -4,12 +4,12 @@ package storetest import ( - "context" "testing" "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -20,6 +20,8 @@ func TestUserAccessTokenStore(t *testing.T, ss store.Store) { } func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { + c := request.TestContext(t) + uat := &model.UserAccessToken{ Token: model.NewId(), UserId: model.NewId(), @@ -30,7 +32,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { s1.UserId = uat.UserId s1.Token = uat.Token - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) _, nErr := ss.UserAccessToken().Save(uat) @@ -58,7 +60,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { nErr = ss.UserAccessToken().Delete(uat.Id) require.NoError(t, nErr) - _, err = ss.Session().Get(context.Background(), s1.Token) + _, err = ss.Session().Get(c, s1.Token) require.Error(t, err, "should error - session should be deleted") _, nErr = ss.UserAccessToken().GetByToken(s1.Token) @@ -68,7 +70,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { s2.UserId = uat.UserId s2.Token = uat.Token - s2, err = ss.Session().Save(s2) + s2, err = ss.Session().Save(c, s2) require.NoError(t, err) _, nErr = ss.UserAccessToken().Save(uat) @@ -77,7 +79,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { nErr = ss.UserAccessToken().DeleteAllForUser(uat.UserId) require.NoError(t, nErr) - _, err = ss.Session().Get(context.Background(), s2.Token) + _, err = ss.Session().Get(c, s2.Token) require.Error(t, err, "should error - session should be deleted") _, nErr = ss.UserAccessToken().GetByToken(s2.Token) @@ -85,6 +87,8 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { } func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) { + c := request.TestContext(t) + uat := &model.UserAccessToken{ Token: model.NewId(), UserId: model.NewId(), @@ -95,7 +99,7 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) { s1.UserId = uat.UserId s1.Token = uat.Token - s1, err := ss.Session().Save(s1) + s1, err := ss.Session().Save(c, s1) require.NoError(t, err) _, nErr := ss.UserAccessToken().Save(uat) @@ -104,14 +108,14 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) { nErr = ss.UserAccessToken().UpdateTokenDisable(uat.Id) require.NoError(t, nErr) - _, err = ss.Session().Get(context.Background(), s1.Token) + _, err = ss.Session().Get(c, s1.Token) require.Error(t, err, "should error - session should be deleted") s2 := &model.Session{} s2.UserId = uat.UserId s2.Token = uat.Token - _, err = ss.Session().Save(s2) + _, err = ss.Session().Save(c, s2) require.NoError(t, err) nErr = ss.UserAccessToken().UpdateTokenEnable(uat.Id) @@ -119,6 +123,8 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) { } func testUserAccessTokenSearch(t *testing.T, ss store.Store) { + c := request.TestContext(t) + u1 := model.User{} u1.Email = MakeEmail() u1.Username = model.NewId() @@ -136,7 +142,7 @@ func testUserAccessTokenSearch(t *testing.T, ss store.Store) { s1.UserId = uat.UserId s1.Token = uat.Token - _, nErr := ss.Session().Save(s1) + _, nErr := ss.Session().Save(c, s1) require.NoError(t, nErr) _, nErr = ss.UserAccessToken().Save(uat) diff --git a/server/channels/store/storetest/user_store.go b/server/channels/store/storetest/user_store.go index 6b89cebe15..4a1426aa9a 100644 --- a/server/channels/store/storetest/user_store.go +++ b/server/channels/store/storetest/user_store.go @@ -14,6 +14,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/public/shared/request" "github.com/mattermost/mattermost/server/v8/channels/store" ) @@ -5244,6 +5245,8 @@ func testUserStoreGetChannelGroupUsers(t *testing.T, ss store.Store) { } func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { + c := request.TestContext(t) + // create users t.Run("Must do nothing with regular user", func(t *testing.T) { id := model.NewId() @@ -5280,7 +5283,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.Equal(t, "system_user", updatedUser.Roles) require.True(t, user.UpdateAt < updatedUser.UpdateAt) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5325,7 +5328,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user system_admin", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5381,7 +5384,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5421,7 +5424,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5466,7 +5469,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user custom_role", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5532,7 +5535,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId1, user1.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId1, user1.Id) require.NoError(t, nErr) require.False(t, updatedTeamMember.SchemeGuest) require.True(t, updatedTeamMember.SchemeUser) @@ -5546,7 +5549,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", notUpdatedUser.Roles) - notUpdatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId2, user2.Id) + notUpdatedTeamMember, nErr := ss.Team().GetMember(c, teamId2, user2.Id) require.NoError(t, nErr) require.True(t, notUpdatedTeamMember.SchemeGuest) require.False(t, notUpdatedTeamMember.SchemeUser) @@ -5559,6 +5562,8 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) { } func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { + c := request.TestContext(t) + // create users t.Run("Must do nothing with guest", func(t *testing.T) { id := model.NewId() @@ -5593,7 +5598,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.Equal(t, "system_guest", updatedUser.Roles) require.True(t, user.UpdateAt < updatedUser.UpdateAt) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, updatedUser.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, updatedUser.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5636,7 +5641,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5688,7 +5693,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5726,7 +5731,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5769,7 +5774,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5833,7 +5838,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_guest", updatedUser.Roles) - updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId1, user1.Id) + updatedTeamMember, nErr := ss.Team().GetMember(c, teamId1, user1.Id) require.NoError(t, nErr) require.True(t, updatedTeamMember.SchemeGuest) require.False(t, updatedTeamMember.SchemeUser) @@ -5847,7 +5852,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) { require.NoError(t, err) require.Equal(t, "system_user", notUpdatedUser.Roles) - notUpdatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId2, user2.Id) + notUpdatedTeamMember, nErr := ss.Team().GetMember(c, teamId2, user2.Id) require.NoError(t, nErr) require.False(t, notUpdatedTeamMember.SchemeGuest) require.True(t, notUpdatedTeamMember.SchemeUser) diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index 2b99cb7816..2c2a20cdaf 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -779,10 +779,10 @@ func (s *TimerLayerChannelStore) CreateDirectChannel(userID *model.User, otherUs return result, err } -func (s *TimerLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { +func (s *TimerLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) { start := time.Now() - result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts) + result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -2947,10 +2947,10 @@ func (s *TimerLayerComplianceStore) GetAll(offset int, limit int) (model.Complia return result, err } -func (s *TimerLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { +func (s *TimerLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) { start := time.Now() - result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit) + result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -3155,10 +3155,10 @@ func (s *TimerLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64) error return err } -func (s *TimerLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { +func (s *TimerLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) { start := time.Now() - result, err := s.EmojiStore.Get(ctx, id, allowFromCache) + result, err := s.EmojiStore.Get(c, id, allowFromCache) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -3171,10 +3171,10 @@ func (s *TimerLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bo return result, err } -func (s *TimerLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { +func (s *TimerLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) { start := time.Now() - result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache) + result, err := s.EmojiStore.GetByName(c, name, allowFromCache) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -3203,10 +3203,10 @@ func (s *TimerLayerEmojiStore) GetList(offset int, limit int, sort string) ([]*m return result, err } -func (s *TimerLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) { +func (s *TimerLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) { start := time.Now() - result, err := s.EmojiStore.GetMultipleByName(ctx, names) + result, err := s.EmojiStore.GetMultipleByName(c, names) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -4721,10 +4721,10 @@ func (s *TimerLayerJobStore) UpdateStatusOptimistically(id string, currentStatus return result, err } -func (s *TimerLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) { +func (s *TimerLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) { start := time.Now() - result, err := s.LicenseStore.Get(ctx, id) + result, err := s.LicenseStore.Get(c, id) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -7471,10 +7471,10 @@ func (s *TimerLayerSessionStore) Cleanup(expiryTime int64, batchSize int64) erro return err } -func (s *TimerLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) { +func (s *TimerLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) { start := time.Now() - result, err := s.SessionStore.Get(ctx, sessionIDOrToken) + result, err := s.SessionStore.Get(c, sessionIDOrToken) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -7487,10 +7487,10 @@ func (s *TimerLayerSessionStore) Get(ctx context.Context, sessionIDOrToken strin return result, err } -func (s *TimerLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) { +func (s *TimerLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) { start := time.Now() - result, err := s.SessionStore.GetSessions(userID) + result, err := s.SessionStore.GetSessions(c, userID) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -7583,10 +7583,10 @@ func (s *TimerLayerSessionStore) RemoveAllSessions() error { return err } -func (s *TimerLayerSessionStore) Save(session *model.Session) (*model.Session, error) { +func (s *TimerLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) { start := time.Now() - result, err := s.SessionStore.Save(session) + result, err := s.SessionStore.Save(c, session) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -8670,10 +8670,10 @@ func (s *TimerLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) { return result, err } -func (s *TimerLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) { +func (s *TimerLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) { start := time.Now() - result, err := s.TeamStore.GetMember(ctx, teamID, userID) + result, err := s.TeamStore.GetMember(c, teamID, userID) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -8766,10 +8766,10 @@ func (s *TimerLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Team, er return result, err } -func (s *TimerLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { +func (s *TimerLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) { start := time.Now() - result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted) + result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -9756,10 +9756,10 @@ func (s *TimerLayerUploadSessionStore) Delete(id string) error { return err } -func (s *TimerLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { +func (s *TimerLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) { start := time.Now() - result, err := s.UploadSessionStore.Get(ctx, id) + result, err := s.UploadSessionStore.Get(c, id) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { diff --git a/server/channels/web/handlers_test.go b/server/channels/web/handlers_test.go index 5fd88a288a..8076d9a8dc 100644 --- a/server/channels/web/handlers_test.go +++ b/server/channels/web/handlers_test.go @@ -171,7 +171,7 @@ func TestHandlerServeCSRFToken(t *testing.T) { } session.GenerateCSRF() th.App.SetSessionExpireInHours(session, 24) - session, err := th.App.CreateSession(session) + session, err := th.App.CreateSession(th.Context, session) if err != nil { t.Errorf("Expected nil, got %s", err) } diff --git a/server/channels/web/oauth.go b/server/channels/web/oauth.go index 3323e4cccd..67a9890c27 100644 --- a/server/channels/web/oauth.go +++ b/server/channels/web/oauth.go @@ -69,7 +69,7 @@ func authorizeOAuthApp(c *Context, w http.ResponseWriter, r *http.Request) { defer c.LogAuditRec(auditRec) c.LogAudit("attempt") - redirectURL, appErr := c.App.AllowOAuthAppAccessToUser(c.AppContext.Session().UserId, authRequest) + redirectURL, appErr := c.App.AllowOAuthAppAccessToUser(c.AppContext, c.AppContext.Session().UserId, authRequest) if appErr != nil { c.Err = appErr return @@ -93,7 +93,7 @@ func deauthorizeOAuthApp(c *Context, w http.ResponseWriter, r *http.Request) { auditRec := c.MakeAuditRecord("deauthorizeOAuthApp", audit.Fail) defer c.LogAuditRec(auditRec) - err := c.App.DeauthorizeOAuthAppForUser(c.AppContext.Session().UserId, clientId) + err := c.App.DeauthorizeOAuthAppForUser(c.AppContext, c.AppContext.Session().UserId, clientId) if err != nil { c.Err = err return @@ -166,7 +166,7 @@ func authorizeOAuthPage(c *Context, w http.ResponseWriter, r *http.Request) { // Automatically allow if the app is trusted if oauthApp.IsTrusted || isAuthorized { - redirectURL, err := c.App.AllowOAuthAppAccessToUser(c.AppContext.Session().UserId, authRequest) + redirectURL, err := c.App.AllowOAuthAppAccessToUser(c.AppContext, c.AppContext.Session().UserId, authRequest) if err != nil { utils.RenderWebAppError(c.App.Config(), w, r, err, c.App.AsymmetricSigningKey()) @@ -229,7 +229,7 @@ func getAccessToken(c *Context, w http.ResponseWriter, r *http.Request) { auditRec.AddMeta("client_id", clientId) c.LogAudit("attempt") - accessRsp, err := c.App.GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectURI, code, secret, refreshToken) + accessRsp, err := c.App.GetOAuthAccessTokenForCodeFlow(c.AppContext, clientId, grantType, redirectURI, code, secret, refreshToken) if err != nil { c.Err = err return diff --git a/server/channels/web/oauth_test.go b/server/channels/web/oauth_test.go index a632512d03..28cfbe1321 100644 --- a/server/channels/web/oauth_test.go +++ b/server/channels/web/oauth_test.go @@ -746,7 +746,7 @@ func (th *TestHelper) Login(client *model.Client4, user *model.User) { Roles: user.GetRawRoles(), IsOAuth: false, } - session, _ = th.App.CreateSession(session) + session, _ = th.App.CreateSession(th.Context, session) client.AuthToken = session.Token client.AuthType = model.HeaderBearer } diff --git a/server/channels/web/saml.go b/server/channels/web/saml.go index 139bd124af..798d53e9d0 100644 --- a/server/channels/web/saml.go +++ b/server/channels/web/saml.go @@ -160,7 +160,7 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { c.App.AddDirectChannels(c.AppContext, teamId, user) } case model.OAuthActionEmailToSSO: - if err = c.App.RevokeAllSessions(user.Id); err != nil { + if err = c.App.RevokeAllSessions(c.AppContext, user.Id); err != nil { c.Err = err return } diff --git a/server/cmd/mattermost/commands/export.go b/server/cmd/mattermost/commands/export.go index 816e091c87..042e549672 100644 --- a/server/cmd/mattermost/commands/export.go +++ b/server/cmd/mattermost/commands/export.go @@ -182,7 +182,7 @@ func buildExportCmdF(format string) func(command *cobra.Command, args []string) return errors.New("message export feature not available") } - warningsCount, appErr := a.MessageExport().RunExport(format, startTime, limit) + warningsCount, appErr := a.MessageExport().RunExport(request.EmptyContext(a.Log()), format, startTime, limit) if appErr != nil { return appErr } diff --git a/server/cmd/mmctl/commands/bot_e2e_test.go b/server/cmd/mmctl/commands/bot_e2e_test.go index b771819868..05091052e5 100644 --- a/server/cmd/mmctl/commands/bot_e2e_test.go +++ b/server/cmd/mmctl/commands/bot_e2e_test.go @@ -426,7 +426,7 @@ func (s *MmctlE2ETestSuite) TestBotCreateCmdF() { token, ok := printer.GetLines()[1].(*model.UserAccessToken) s.Require().True(ok) defer func() { - err := s.th.App.RevokeUserAccessToken(token) + err := s.th.App.RevokeUserAccessToken(s.th.Context, token) s.Require().Nil(err) }() s.Require().Empty(printer.GetErrorLines()) diff --git a/server/einterfaces/message_export.go b/server/einterfaces/message_export.go index c341f50d27..32e4cbc552 100644 --- a/server/einterfaces/message_export.go +++ b/server/einterfaces/message_export.go @@ -10,5 +10,5 @@ import ( type MessageExportInterface interface { StartSynchronizeJob(c *request.Context, exportFromTimestamp int64) (*model.Job, *model.AppError) - RunExport(format string, since int64, limit int) (int64, *model.AppError) + RunExport(c *request.Context, format string, since int64, limit int) (int64, *model.AppError) } diff --git a/server/einterfaces/mocks/MessageExportInterface.go b/server/einterfaces/mocks/MessageExportInterface.go index 06571fe105..3aabcbdc55 100644 --- a/server/einterfaces/mocks/MessageExportInterface.go +++ b/server/einterfaces/mocks/MessageExportInterface.go @@ -15,23 +15,23 @@ type MessageExportInterface struct { mock.Mock } -// RunExport provides a mock function with given fields: format, since, limit -func (_m *MessageExportInterface) RunExport(format string, since int64, limit int) (int64, *model.AppError) { - ret := _m.Called(format, since, limit) +// RunExport provides a mock function with given fields: c, format, since, limit +func (_m *MessageExportInterface) RunExport(c *request.Context, format string, since int64, limit int) (int64, *model.AppError) { + ret := _m.Called(c, format, since, limit) var r0 int64 var r1 *model.AppError - if rf, ok := ret.Get(0).(func(string, int64, int) (int64, *model.AppError)); ok { - return rf(format, since, limit) + if rf, ok := ret.Get(0).(func(*request.Context, string, int64, int) (int64, *model.AppError)); ok { + return rf(c, format, since, limit) } - if rf, ok := ret.Get(0).(func(string, int64, int) int64); ok { - r0 = rf(format, since, limit) + if rf, ok := ret.Get(0).(func(*request.Context, string, int64, int) int64); ok { + r0 = rf(c, format, since, limit) } else { r0 = ret.Get(0).(int64) } - if rf, ok := ret.Get(1).(func(string, int64, int) *model.AppError); ok { - r1 = rf(format, since, limit) + if rf, ok := ret.Get(1).(func(*request.Context, string, int64, int) *model.AppError); ok { + r1 = rf(c, format, since, limit) } else { if ret.Get(1) != nil { r1 = ret.Get(1).(*model.AppError) diff --git a/server/public/shared/request/context.go b/server/public/shared/request/context.go index d2d71a1768..3e43611287 100644 --- a/server/public/shared/request/context.go +++ b/server/public/shared/request/context.go @@ -49,7 +49,7 @@ func EmptyContext(logger mlog.LoggerIFace) *Context { // TestContext creates an empty context with a new logger to use in testing where a test helper is // not required. -func TestContext(t *testing.T) *Context { +func TestContext(t testing.TB) *Context { logger := mlog.CreateConsoleTestLogger(t) return EmptyContext(logger) }