MM-30882: Fix read-after-write issue for demoting user (#16911)
* MM-30882: Fix read-after-write issue for demoting user In (*App).DemoteUserToGuest, we would demote a user, and then immediately read it back to do future operations from the user. This reading back of the user had the effect of sticking the old value into the cache after which it would never be updated. There was another issue along with this, which was when the invalidation message would broadcast across the cluster, it would hit the cache invalidation problem where an unrelated store call would miss the cache because it was invalidated, and then again read from replica and stick the old value. To fix all these, we return the new value directly from the store method to avoid having the app to read it again. And we add a map in the localcache layer which tracks invalidations made, and then switch to use master if it's true. The core change is fairly limited, but due to changing the store method signatures, a lot of code needed to be updated to pass "context.Background". Therefore the PR just "appears" to be big, but the main changes are limited to app/user.go, sqlstore/user_store.go and user_layer.go https://mattermost.atlassian.net/browse/MM-30882 ```release-note Fix an issue where demoting a user to guest would not take effect in an environment with read replicas. ``` * Fix concurrent map access * Fixing mistakes * fix tests
Этот коммит содержится в:
коммит произвёл
GitHub
родитель
49907d3081
Коммит
021c90f29f
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
@@ -55,7 +56,7 @@ func (a *App) CreateBot(bot *model.Bot) (*model.Bot, *model.AppError) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Get the owner of the bot, if one exists. If not, don't send a message
|
// Get the owner of the bot, if one exists. If not, don't send a message
|
||||||
ownerUser, err := a.Srv().Store.User().Get(bot.OwnerId)
|
ownerUser, err := a.Srv().Store.User().Get(context.Background(), bot.OwnerId)
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
if err != nil && !errors.As(err, &nfErr) {
|
if err != nil && !errors.As(err, &nfErr) {
|
||||||
return nil, model.NewAppError("CreateBot", "app.user.get.app_error", nil, err.Error(), http.StatusInternalServerError)
|
return nil, model.NewAppError("CreateBot", "app.user.get.app_error", nil, err.Error(), http.StatusInternalServerError)
|
||||||
@@ -156,7 +157,7 @@ func (a *App) PatchBot(botUserId string, botPatch *model.BotPatch) (*model.Bot,
|
|||||||
|
|
||||||
bot.Patch(botPatch)
|
bot.Patch(botPatch)
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(botUserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), botUserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -233,7 +234,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.
|
// UpdateBotActive marks a bot as active or inactive, along with its corresponding user.
|
||||||
func (a *App) UpdateBotActive(botUserId string, active bool) (*model.Bot, *model.AppError) {
|
func (a *App) UpdateBotActive(botUserId string, active bool) (*model.Bot, *model.AppError) {
|
||||||
user, nErr := a.Srv().Store.User().Get(botUserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), botUserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -67,7 +68,7 @@ func (a *App) JoinDefaultChannels(teamID string, user *model.User, shouldBeAdmin
|
|||||||
var requestor *model.User
|
var requestor *model.User
|
||||||
var nErr error
|
var nErr error
|
||||||
if userRequestorId != "" {
|
if userRequestorId != "" {
|
||||||
requestor, nErr = a.Srv().Store.User().Get(userRequestorId)
|
requestor, nErr = a.Srv().Store.User().Get(context.Background(), userRequestorId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -264,7 +265,7 @@ func (a *App) CreateChannel(channel *model.Channel, addMember bool) (*model.Chan
|
|||||||
}
|
}
|
||||||
|
|
||||||
if addMember {
|
if addMember {
|
||||||
user, nErr := a.Srv().Store.User().Get(channel.CreatorId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), channel.CreatorId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -386,7 +387,7 @@ func (a *App) handleCreationEvent(userID, otherUserID string, channel *model.Cha
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) createDirectChannel(userID, otherUserID string) (*model.Channel, *model.AppError) {
|
func (a *App) createDirectChannel(userID, otherUserID string) (*model.Channel, *model.AppError) {
|
||||||
users, err := a.Srv().Store.User().GetMany([]string{userID, otherUserID})
|
users, err := a.Srv().Store.User().GetMany(context.Background(), []string{userID, otherUserID})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, model.NewAppError("CreateDirectChannel", "api.channel.create_direct_channel.invalid_user.app_error", nil, err.Error(), http.StatusBadRequest)
|
return nil, model.NewAppError("CreateDirectChannel", "api.channel.create_direct_channel.invalid_user.app_error", nil, err.Error(), http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
@@ -520,7 +521,7 @@ func (a *App) createGroupChannel(userIDs []string) (*model.Channel, *model.AppEr
|
|||||||
return nil, model.NewAppError("CreateGroupChannel", "api.channel.create_group.bad_size.app_error", nil, "", http.StatusBadRequest)
|
return nil, model.NewAppError("CreateGroupChannel", "api.channel.create_group.bad_size.app_error", nil, "", http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
users, err := a.Srv().Store.User().GetProfileByIds(userIDs, nil, true)
|
users, err := a.Srv().Store.User().GetProfileByIds(context.Background(), userIDs, nil, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, model.NewAppError("createGroupChannel", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
return nil, model.NewAppError("createGroupChannel", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
@@ -599,7 +600,7 @@ func (a *App) GetGroupChannel(userIDs []string) (*model.Channel, *model.AppError
|
|||||||
return nil, model.NewAppError("GetGroupChannel", "api.channel.create_group.bad_size.app_error", nil, "", http.StatusBadRequest)
|
return nil, model.NewAppError("GetGroupChannel", "api.channel.create_group.bad_size.app_error", nil, "", http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
users, err := a.Srv().Store.User().GetProfileByIds(userIDs, nil, true)
|
users, err := a.Srv().Store.User().GetProfileByIds(context.Background(), userIDs, nil, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, model.NewAppError("GetGroupChannel", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
return nil, model.NewAppError("GetGroupChannel", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
@@ -745,7 +746,7 @@ func (a *App) RestoreChannel(channel *model.Channel, userID string) (*model.Chan
|
|||||||
message.Add("channel_id", channel.Id)
|
message.Add("channel_id", channel.Id)
|
||||||
a.Publish(message)
|
a.Publish(message)
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(userID)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -1236,7 +1237,7 @@ func (a *App) DeleteChannel(channel *model.Channel, userID string) *model.AppErr
|
|||||||
var user *model.User
|
var user *model.User
|
||||||
if userID != "" {
|
if userID != "" {
|
||||||
var nErr error
|
var nErr error
|
||||||
user, nErr = a.Srv().Store.User().Get(userID)
|
user, nErr = a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -1492,7 +1493,7 @@ func (a *App) AddDirectChannels(teamID string, user *model.User) *model.AppError
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) PostUpdateChannelHeaderMessage(userID string, channel *model.Channel, oldChannelHeader, newChannelHeader string) *model.AppError {
|
func (a *App) PostUpdateChannelHeaderMessage(userID string, channel *model.Channel, oldChannelHeader, newChannelHeader string) *model.AppError {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return model.NewAppError("PostUpdateChannelHeaderMessage", "api.channel.post_update_channel_header_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
return model.NewAppError("PostUpdateChannelHeaderMessage", "api.channel.post_update_channel_header_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
@@ -1526,7 +1527,7 @@ func (a *App) PostUpdateChannelHeaderMessage(userID string, channel *model.Chann
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) PostUpdateChannelPurposeMessage(userID string, channel *model.Channel, oldChannelPurpose string, newChannelPurpose string) *model.AppError {
|
func (a *App) PostUpdateChannelPurposeMessage(userID string, channel *model.Channel, oldChannelPurpose string, newChannelPurpose string) *model.AppError {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return model.NewAppError("PostUpdateChannelPurposeMessage", "app.channel.post_update_channel_purpose_message.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
return model.NewAppError("PostUpdateChannelPurposeMessage", "app.channel.post_update_channel_purpose_message.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
@@ -1559,7 +1560,7 @@ func (a *App) PostUpdateChannelPurposeMessage(userID string, channel *model.Chan
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) PostUpdateChannelDisplayNameMessage(userID string, channel *model.Channel, oldChannelDisplayName, newChannelDisplayName string) *model.AppError {
|
func (a *App) PostUpdateChannelDisplayNameMessage(userID string, channel *model.Channel, oldChannelDisplayName, newChannelDisplayName string) *model.AppError {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return model.NewAppError("PostUpdateChannelDisplayNameMessage", "api.channel.post_update_channel_displayname_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
return model.NewAppError("PostUpdateChannelDisplayNameMessage", "api.channel.post_update_channel_displayname_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
@@ -1905,7 +1906,7 @@ func (a *App) JoinChannel(channel *model.Channel, userID string) *model.AppError
|
|||||||
userChan := make(chan store.StoreResult, 1)
|
userChan := make(chan store.StoreResult, 1)
|
||||||
memberChan := make(chan store.StoreResult, 1)
|
memberChan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
userChan <- store.StoreResult{Data: user, NErr: err}
|
userChan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(userChan)
|
close(userChan)
|
||||||
}()
|
}()
|
||||||
@@ -2014,7 +2015,7 @@ func (a *App) LeaveChannel(channelId string, userID string) *model.AppError {
|
|||||||
|
|
||||||
uc := make(chan store.StoreResult, 1)
|
uc := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
uc <- store.StoreResult{Data: user, NErr: err}
|
uc <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uc)
|
close(uc)
|
||||||
}()
|
}()
|
||||||
@@ -2176,7 +2177,7 @@ func (a *App) postRemoveFromChannelMessage(removerUserId string, removedUser *mo
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) removeUserFromChannel(userIDToRemove string, removerUserId string, channel *model.Channel) *model.AppError {
|
func (a *App) removeUserFromChannel(userIDToRemove string, removerUserId string, channel *model.Channel) *model.AppError {
|
||||||
user, nErr := a.Srv().Store.User().Get(userIDToRemove)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), userIDToRemove)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -1910,7 +1911,7 @@ func TestMarkChannelsAsViewedPanic(t *testing.T) {
|
|||||||
|
|
||||||
mockStore := th.App.Srv().Store.(*mocks.Store)
|
mockStore := th.App.Srv().Store.(*mocks.Store)
|
||||||
mockUserStore := mocks.UserStore{}
|
mockUserStore := mocks.UserStore{}
|
||||||
mockUserStore.On("Get", "userID").Return(nil, model.NewAppError("SqlUserStore.Get", "app.user.get.app_error", nil, "user_id=userID", http.StatusInternalServerError))
|
mockUserStore.On("Get", context.Background(), "userID").Return(nil, model.NewAppError("SqlUserStore.Get", "app.user.get.app_error", nil, "user_id=userID", http.StatusInternalServerError))
|
||||||
mockChannelStore := mocks.ChannelStore{}
|
mockChannelStore := mocks.ChannelStore{}
|
||||||
mockChannelStore.On("Get", "channelID", true).Return(&model.Channel{}, nil)
|
mockChannelStore.On("Get", "channelID", true).Return(&model.Channel{}, nil)
|
||||||
mockChannelStore.On("GetMember", "channelID", "userID").Return(&model.ChannelMember{
|
mockChannelStore.On("GetMember", "channelID", "userID").Return(&model.ChannelMember{
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"io/ioutil"
|
"io/ioutil"
|
||||||
@@ -368,7 +369,7 @@ func (a *App) tryExecuteCustomCommand(args *model.CommandArgs, trigger string, m
|
|||||||
|
|
||||||
userChan := make(chan store.StoreResult, 1)
|
userChan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(args.UserId)
|
user, err := a.Srv().Store.User().Get(context.Background(), args.UserId)
|
||||||
userChan <- store.StoreResult{Data: user, NErr: err}
|
userChan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(userChan)
|
close(userChan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"html/template"
|
"html/template"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -194,7 +195,7 @@ func (job *EmailBatchingJob) checkPendingNotifications(now time.Time, handler fu
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (es *EmailService) sendBatchedEmailNotification(userID string, notifications []*batchedNotification) {
|
func (es *EmailService) sendBatchedEmailNotification(userID string, notifications []*batchedNotification) {
|
||||||
user, err := es.srv.Store.User().Get(userID)
|
user, err := es.srv.Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mlog.Warn("Unable to find recipient for batched email notification")
|
mlog.Warn("Unable to find recipient for batched email notification")
|
||||||
return
|
return
|
||||||
@@ -205,7 +206,7 @@ func (es *EmailService) sendBatchedEmailNotification(userID string, notification
|
|||||||
|
|
||||||
var contents string
|
var contents string
|
||||||
for _, notification := range notifications {
|
for _, notification := range notifications {
|
||||||
sender, err := es.srv.Store.User().Get(notification.post.UserId)
|
sender, err := es.srv.Store.User().Get(context.Background(), notification.post.UserId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mlog.Warn("Unable to find sender of post for batched email notification")
|
mlog.Warn("Unable to find sender of post for batched email notification")
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package app
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"archive/zip"
|
"archive/zip"
|
||||||
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -483,7 +484,7 @@ func (a *App) BuildPostReactions(postId string) (*[]ReactionImportData, *model.A
|
|||||||
}
|
}
|
||||||
|
|
||||||
for _, reaction := range reactions {
|
for _, reaction := range reactions {
|
||||||
user, err := a.Srv().Store.User().Get(reaction.UserId)
|
user, err := a.Srv().Store.User().Get(context.Background(), reaction.UserId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
if errors.As(err, &nfErr) { // this is a valid case, the user that reacted might've been deleted by now
|
if errors.As(err, &nfErr) { // this is a valid case, the user that reacted might've been deleted by now
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ package app
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -84,7 +85,7 @@ func (a *App) DoPostActionWithCookie(postId, actionId, userID, selectedOption st
|
|||||||
|
|
||||||
userChan := make(chan store.StoreResult, 1)
|
userChan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(upstreamRequest.UserId)
|
user, err := a.Srv().Store.User().Get(context.Background(), upstreamRequest.UserId)
|
||||||
userChan <- store.StoreResult{Data: user, NErr: err}
|
userChan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(userChan)
|
close(userChan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
"sort"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -28,7 +29,7 @@ func (a *App) SendNotifications(post *model.Post, team *model.Team, channel *mod
|
|||||||
|
|
||||||
pchan := make(chan store.StoreResult, 1)
|
pchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
props, err := a.Srv().Store.User().GetAllProfilesInChannel(channel.Id, true)
|
props, err := a.Srv().Store.User().GetAllProfilesInChannel(context.Background(), channel.Id, true)
|
||||||
pchan <- store.StoreResult{Data: props, NErr: err}
|
pchan <- store.StoreResult{Data: props, NErr: err}
|
||||||
close(pchan)
|
close(pchan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package app
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
b64 "encoding/base64"
|
b64 "encoding/base64"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -288,7 +289,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectUri, c
|
|||||||
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.redirect_uri.app_error", nil, "", http.StatusBadRequest)
|
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.redirect_uri.app_error", nil, "", http.StatusBadRequest)
|
||||||
}
|
}
|
||||||
|
|
||||||
user, nErr = a.Srv().Store.User().Get(authData.UserId)
|
user, nErr = a.Srv().Store.User().Get(context.Background(), authData.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound)
|
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound)
|
||||||
}
|
}
|
||||||
@@ -347,7 +348,7 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectUri, c
|
|||||||
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.refresh_token.app_error", nil, "", http.StatusNotFound)
|
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.refresh_token.app_error", nil, "", http.StatusNotFound)
|
||||||
}
|
}
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(accessData.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), accessData.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound)
|
return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -53,7 +54,7 @@ func (a *App) CreatePostAsUser(post *model.Post, currentSessionId string, setOnl
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err.Id == "api.post.create_post.town_square_read_only" {
|
if err.Id == "api.post.create_post.town_square_read_only" {
|
||||||
user, nErr := a.Srv().Store.User().Get(post.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), post.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -191,7 +192,7 @@ func (a *App) CreatePost(post *model.Post, channel *model.Channel, triggerWebhoo
|
|||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(post.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), post.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
@@ -427,7 +428,7 @@ func (a *App) SetSessionExpireInDays(session *model.Session, days int) {
|
|||||||
|
|
||||||
func (a *App) CreateUserAccessToken(token *model.UserAccessToken) (*model.UserAccessToken, *model.AppError) {
|
func (a *App) CreateUserAccessToken(token *model.UserAccessToken) (*model.UserAccessToken, *model.AppError) {
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(token.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), token.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -476,7 +477,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)
|
return nil, model.NewAppError("createSessionForUserAccessToken", "app.user_access_token.invalid_or_missing", nil, "inactive_token", http.StatusUnauthorized)
|
||||||
}
|
}
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(token.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), token.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package app
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"mime/multipart"
|
"mime/multipart"
|
||||||
"regexp"
|
"regexp"
|
||||||
@@ -88,7 +89,7 @@ func replaceUserIds(userStore store.UserStore, text string) string {
|
|||||||
userIDs = append(userIDs, match[1])
|
userIDs = append(userIDs, match[1])
|
||||||
}
|
}
|
||||||
|
|
||||||
if users, err := userStore.GetProfileByIds(userIDs, nil, true); err == nil {
|
if users, err := userStore.GetProfileByIds(context.Background(), userIDs, nil, true); err == nil {
|
||||||
for _, user := range users {
|
for _, user := range users {
|
||||||
text = strings.Replace(text, "<@"+user.Id+">", "@"+user.Username, -1)
|
text = strings.Replace(text, "<@"+user.Id+">", "@"+user.Username, -1)
|
||||||
}
|
}
|
||||||
|
|||||||
14
app/team.go
14
app/team.go
@@ -489,7 +489,7 @@ func (a *App) AddUserToTeam(teamID string, userID string, userRequestorId string
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
@@ -566,7 +566,7 @@ func (a *App) AddUserToTeamByToken(userID string, tokenID string) (*model.Team,
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
@@ -641,7 +641,7 @@ func (a *App) AddUserToTeamByInviteId(inviteId string, userID string) (*model.Te
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
@@ -1163,7 +1163,7 @@ func (a *App) RemoveUserFromTeam(teamID string, userID string, requestorId strin
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
@@ -1206,7 +1206,7 @@ func (a *App) RemoveTeamMemberFromTeam(teamMember *model.TeamMember, requestorId
|
|||||||
message.Add("team_id", teamMember.TeamId)
|
message.Add("team_id", teamMember.TeamId)
|
||||||
a.Publish(message)
|
a.Publish(message)
|
||||||
|
|
||||||
user, nErr := a.Srv().Store.User().Get(teamMember.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), teamMember.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -1368,7 +1368,7 @@ func (a *App) prepareInviteNewUsersToTeam(teamID, senderId string) (*model.User,
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(senderId)
|
user, err := a.Srv().Store.User().Get(context.Background(), senderId)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
@@ -1499,7 +1499,7 @@ func (a *App) prepareInviteGuestsToChannels(teamID string, guestsInvite *model.G
|
|||||||
}()
|
}()
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(senderId)
|
user, err := a.Srv().Store.User().Get(context.Background(), senderId)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
26
app/user.go
26
app/user.go
@@ -5,6 +5,7 @@ package app
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
b64 "encoding/base64"
|
b64 "encoding/base64"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -432,7 +433,7 @@ func (a *App) IsUsernameTaken(name string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) GetUser(userID string) (*model.User, *model.AppError) {
|
func (a *App) GetUser(userID string) (*model.User, *model.AppError) {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -683,7 +684,7 @@ func (a *App) GetChannelGroupUsers(channelID string) ([]*model.User, *model.AppE
|
|||||||
func (a *App) GetUsersByIds(userIDs []string, options *store.UserGetByIdsOpts) ([]*model.User, *model.AppError) {
|
func (a *App) GetUsersByIds(userIDs []string, options *store.UserGetByIdsOpts) ([]*model.User, *model.AppError) {
|
||||||
allowFromCache := options.ViewRestrictions == nil
|
allowFromCache := options.ViewRestrictions == nil
|
||||||
|
|
||||||
users, err := a.Srv().Store.User().GetProfileByIds(userIDs, options, allowFromCache)
|
users, err := a.Srv().Store.User().GetProfileByIds(context.Background(), userIDs, options, allowFromCache)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, model.NewAppError("GetUsersByIds", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
return nil, model.NewAppError("GetUsersByIds", "app.user.get_profiles.app_error", nil, err.Error(), http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
@@ -739,7 +740,7 @@ func (a *App) GenerateMfaSecret(userID string) (*model.MfaSecret, *model.AppErro
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) ActivateMfa(userID, token string) *model.AppError {
|
func (a *App) ActivateMfa(userID, token string) *model.AppError {
|
||||||
user, err := a.Srv().Store.User().Get(userID)
|
user, err := a.Srv().Store.User().Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -1212,7 +1213,7 @@ func (a *App) sendUpdatedUserEvent(user model.User) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) UpdateUser(user *model.User, sendNotifications bool) (*model.User, *model.AppError) {
|
func (a *App) UpdateUser(user *model.User, sendNotifications bool) (*model.User, *model.AppError) {
|
||||||
prev, err := a.Srv().Store.User().Get(user.Id)
|
prev, err := a.Srv().Store.User().Get(context.Background(), user.Id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
@@ -2058,7 +2059,7 @@ func (a *App) FilterNonGroupChannelMembers(userIDs []string, channel *model.Chan
|
|||||||
// and returns the list of normal users present in userIDs but not in groupUsers.
|
// and returns the list of normal users present in userIDs but not in groupUsers.
|
||||||
func (a *App) filterNonGroupUsers(userIDs []string, groupUsers []*model.User) ([]string, error) {
|
func (a *App) filterNonGroupUsers(userIDs []string, groupUsers []*model.User) ([]string, error) {
|
||||||
nonMemberIds := []string{}
|
nonMemberIds := []string{}
|
||||||
users, err := a.Srv().Store.User().GetProfileByIds(userIDs, nil, false)
|
users, err := a.Srv().Store.User().GetProfileByIds(context.Background(), userIDs, nil, false)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -2223,19 +2224,14 @@ func (a *App) PromoteGuestToUser(user *model.User, requestorId string) *model.Ap
|
|||||||
// DemoteUserToGuest Convert user's roles and all his mermbership's roles from
|
// DemoteUserToGuest Convert user's roles and all his mermbership's roles from
|
||||||
// regular user roles to guest roles.
|
// regular user roles to guest roles.
|
||||||
func (a *App) DemoteUserToGuest(user *model.User) *model.AppError {
|
func (a *App) DemoteUserToGuest(user *model.User) *model.AppError {
|
||||||
nErr := a.Srv().Store.User().DemoteUserToGuest(user.Id)
|
demotedUser, nErr := a.Srv().Store.User().DemoteUserToGuest(user.Id)
|
||||||
a.InvalidateCacheForUser(user.Id)
|
a.InvalidateCacheForUser(user.Id)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
return model.NewAppError("DemoteUserToGuest", "app.user.demote_user_to_guest.user_update.app_error", nil, nErr.Error(), http.StatusInternalServerError)
|
return model.NewAppError("DemoteUserToGuest", "app.user.demote_user_to_guest.user_update.app_error", nil, nErr.Error(), http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
|
|
||||||
demotedUser, err := a.GetUser(user.Id)
|
a.sendUpdatedUserEvent(*demotedUser)
|
||||||
if err != nil {
|
a.UpdateSessionsIsGuest(demotedUser.Id, demotedUser.IsGuest())
|
||||||
mlog.Warn("Failed to get user on demote user to guest", mlog.Err(err))
|
|
||||||
} else {
|
|
||||||
a.sendUpdatedUserEvent(*demotedUser)
|
|
||||||
a.UpdateSessionsIsGuest(demotedUser.Id, demotedUser.IsGuest())
|
|
||||||
}
|
|
||||||
|
|
||||||
teamMembers, err := a.GetTeamMembersForUser(user.Id)
|
teamMembers, err := a.GetTeamMembersForUser(user.Id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -2248,6 +2244,7 @@ func (a *App) DemoteUserToGuest(user *model.User) *model.AppError {
|
|||||||
channelMembers, err := a.GetChannelMembersForUser(member.TeamId, user.Id)
|
channelMembers, err := a.GetChannelMembersForUser(member.TeamId, user.Id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
mlog.Warn("Failed to get channel members for users on demote user to guest", mlog.Err(err))
|
mlog.Warn("Failed to get channel members for users on demote user to guest", mlog.Err(err))
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, member := range *channelMembers {
|
for _, member := range *channelMembers {
|
||||||
@@ -2260,7 +2257,6 @@ func (a *App) DemoteUserToGuest(user *model.User) *model.AppError {
|
|||||||
}
|
}
|
||||||
|
|
||||||
a.ClearSessionCacheForUser(user.Id)
|
a.ClearSessionCacheForUser(user.Id)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2308,7 +2304,7 @@ func (a *App) GetKnownUsers(userID string) ([]string, *model.AppError) {
|
|||||||
|
|
||||||
// ConvertBotToUser converts a bot to user.
|
// ConvertBotToUser converts a bot to user.
|
||||||
func (a *App) ConvertBotToUser(bot *model.Bot, userPatch *model.UserPatch, sysadmin bool) (*model.User, *model.AppError) {
|
func (a *App) ConvertBotToUser(bot *model.Bot, userPatch *model.UserPatch, sysadmin bool) (*model.User, *model.AppError) {
|
||||||
user, nErr := a.Srv().Store.User().Get(bot.UserId)
|
user, nErr := a.Srv().Store.User().Get(context.Background(), bot.UserId)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
var nfErr *store.ErrNotFound
|
var nfErr *store.ErrNotFound
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package app
|
package app
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -674,7 +675,7 @@ func (a *App) HandleIncomingWebhook(hookID string, req *model.IncomingWebhookReq
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
user, err := a.Srv().Store.User().Get(hook.UserId)
|
user, err := a.Srv().Store.User().Get(context.Background(), hook.UserId)
|
||||||
uchan <- store.StoreResult{Data: user, NErr: err}
|
uchan <- store.StoreResult{Data: user, NErr: err}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package commands
|
package commands
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -115,7 +116,7 @@ func TestDeleteUserBotUser(t *testing.T) {
|
|||||||
defer th.TearDown()
|
defer th.TearDown()
|
||||||
|
|
||||||
th.CheckCommand(t, "user", "delete", th.BasicUser.Username, "--confirm")
|
th.CheckCommand(t, "user", "delete", th.BasicUser.Username, "--confirm")
|
||||||
_, err := th.App.Srv().Store.User().Get(th.BasicUser.Id)
|
_, err := th.App.Srv().Store.User().Get(context.Background(), th.BasicUser.Id)
|
||||||
require.Error(t, err)
|
require.Error(t, err)
|
||||||
|
|
||||||
// Make a bot
|
// Make a bot
|
||||||
@@ -131,7 +132,7 @@ func TestDeleteUserBotUser(t *testing.T) {
|
|||||||
require.Nil(t, nErr)
|
require.Nil(t, nErr)
|
||||||
|
|
||||||
th.CheckCommand(t, "user", "delete", bot.Username, "--confirm")
|
th.CheckCommand(t, "user", "delete", bot.Username, "--confirm")
|
||||||
_, err = th.App.Srv().Store.User().Get(user.Id)
|
_, err = th.App.Srv().Store.User().Get(context.Background(), user.Id)
|
||||||
require.Error(t, err)
|
require.Error(t, err)
|
||||||
_, nErr = th.App.Srv().Store.Bot().Get(user.Id, true)
|
_, nErr = th.App.Srv().Store.Bot().Get(user.Id, true)
|
||||||
require.Error(t, nErr)
|
require.Error(t, nErr)
|
||||||
@@ -199,7 +200,7 @@ func TestConvertUser(t *testing.T) {
|
|||||||
_, err = th.App.Srv().Store.Bot().Get(th.BasicUser2.Id, false)
|
_, err = th.App.Srv().Store.Bot().Get(th.BasicUser2.Id, false)
|
||||||
require.NotNil(t, err)
|
require.NotNil(t, err)
|
||||||
|
|
||||||
user, appErr := th.App.Srv().Store.User().Get(th.BasicUser2.Id)
|
user, appErr := th.App.Srv().Store.User().Get(context.Background(), th.BasicUser2.Id)
|
||||||
require.Nil(t, appErr)
|
require.Nil(t, appErr)
|
||||||
require.Equal(t, "newusername", user.Username)
|
require.Equal(t, "newusername", user.Username)
|
||||||
require.Equal(t, "valid@email.com", user.Email)
|
require.Equal(t, "valid@email.com", user.Email)
|
||||||
|
|||||||
@@ -4,6 +4,8 @@
|
|||||||
package commands
|
package commands
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
"github.com/mattermost/mattermost-server/v5/app"
|
"github.com/mattermost/mattermost-server/v5/app"
|
||||||
"github.com/mattermost/mattermost-server/v5/model"
|
"github.com/mattermost/mattermost-server/v5/model"
|
||||||
)
|
)
|
||||||
@@ -28,7 +30,7 @@ func getUserFromUserArg(a *app.App, userArg string) *model.User {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if user == nil {
|
if user == nil {
|
||||||
user, _ = a.Srv().Store.User().Get(userArg)
|
user, _ = a.Srv().Store.User().Get(context.Background(), userArg)
|
||||||
}
|
}
|
||||||
|
|
||||||
return user
|
return user
|
||||||
|
|||||||
@@ -98,7 +98,7 @@ type LocalCacheStore struct {
|
|||||||
postLastPostsCache cache.Cache
|
postLastPostsCache cache.Cache
|
||||||
lastPostTimeCache cache.Cache
|
lastPostTimeCache cache.Cache
|
||||||
|
|
||||||
user LocalCacheUserStore
|
user *LocalCacheUserStore
|
||||||
userProfileByIdsCache cache.Cache
|
userProfileByIdsCache cache.Cache
|
||||||
profilesInChannelCache cache.Cache
|
profilesInChannelCache cache.Cache
|
||||||
|
|
||||||
@@ -283,7 +283,12 @@ func NewLocalCacheLayer(baseStore store.Store, metrics einterfaces.MetricsInterf
|
|||||||
}); err != nil {
|
}); err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
localCacheStore.user = LocalCacheUserStore{UserStore: baseStore.User(), rootStore: &localCacheStore}
|
localCacheStore.user = &LocalCacheUserStore{
|
||||||
|
UserStore: baseStore.User(),
|
||||||
|
rootStore: &localCacheStore,
|
||||||
|
userProfileByIdsInvalidations: make(map[string]bool),
|
||||||
|
profilesInChannelInvalidations: make(map[string]bool),
|
||||||
|
}
|
||||||
|
|
||||||
// Teams
|
// Teams
|
||||||
if localCacheStore.teamAllTeamIdsForUserCache, err = cacheProvider.NewCache(&cache.CacheOptions{
|
if localCacheStore.teamAllTeamIdsForUserCache, err = cacheProvider.NewCache(&cache.CacheOptions{
|
||||||
|
|||||||
@@ -132,16 +132,16 @@ func getMockStore() *mocks.Store {
|
|||||||
AuthService: "authService",
|
AuthService: "authService",
|
||||||
}}
|
}}
|
||||||
mockUserStore := mocks.UserStore{}
|
mockUserStore := mocks.UserStore{}
|
||||||
mockUserStore.On("GetProfileByIds", []string{"123"}, &store.UserGetByIdsOpts{}, true).Return(fakeUser, nil)
|
mockUserStore.On("GetProfileByIds", mock.Anything, []string{"123"}, &store.UserGetByIdsOpts{}, true).Return(fakeUser, nil)
|
||||||
mockUserStore.On("GetProfileByIds", []string{"123"}, &store.UserGetByIdsOpts{}, false).Return(fakeUser, nil)
|
mockUserStore.On("GetProfileByIds", mock.Anything, []string{"123"}, &store.UserGetByIdsOpts{}, false).Return(fakeUser, nil)
|
||||||
|
|
||||||
fakeProfilesInChannelMap := map[string]*model.User{
|
fakeProfilesInChannelMap := map[string]*model.User{
|
||||||
"456": {Id: "456"},
|
"456": {Id: "456"},
|
||||||
}
|
}
|
||||||
mockUserStore.On("GetAllProfilesInChannel", "123", true).Return(fakeProfilesInChannelMap, nil)
|
mockUserStore.On("GetAllProfilesInChannel", mock.Anything, "123", true).Return(fakeProfilesInChannelMap, nil)
|
||||||
mockUserStore.On("GetAllProfilesInChannel", "123", false).Return(fakeProfilesInChannelMap, nil)
|
mockUserStore.On("GetAllProfilesInChannel", mock.Anything, "123", false).Return(fakeProfilesInChannelMap, nil)
|
||||||
|
|
||||||
mockUserStore.On("Get", "123").Return(fakeUser[0], nil)
|
mockUserStore.On("Get", mock.Anything, "123").Return(fakeUser[0], nil)
|
||||||
users := []*model.User{
|
users := []*model.User{
|
||||||
fakeUser[0],
|
fakeUser[0],
|
||||||
{
|
{
|
||||||
@@ -150,8 +150,8 @@ func getMockStore() *mocks.Store {
|
|||||||
AuthService: "authService",
|
AuthService: "authService",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
mockUserStore.On("GetMany", []string{"123", "456"}).Return(users, nil)
|
mockUserStore.On("GetMany", mock.Anything, []string{"123", "456"}).Return(users, nil)
|
||||||
mockUserStore.On("GetMany", []string{"123"}).Return(users[0:1], nil)
|
mockUserStore.On("GetMany", mock.Anything, []string{"123"}).Return(users[0:1], nil)
|
||||||
mockStore.On("User").Return(&mockUserStore)
|
mockStore.On("User").Return(&mockUserStore)
|
||||||
|
|
||||||
fakeUserTeamIds := []string{"1", "2", "3"}
|
fakeUserTeamIds := []string{"1", "2", "3"}
|
||||||
|
|||||||
@@ -4,21 +4,31 @@
|
|||||||
package localcachelayer
|
package localcachelayer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"sort"
|
"sort"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/mattermost/mattermost-server/v5/model"
|
"github.com/mattermost/mattermost-server/v5/model"
|
||||||
"github.com/mattermost/mattermost-server/v5/store"
|
"github.com/mattermost/mattermost-server/v5/store"
|
||||||
|
"github.com/mattermost/mattermost-server/v5/store/sqlstore"
|
||||||
)
|
)
|
||||||
|
|
||||||
type LocalCacheUserStore struct {
|
type LocalCacheUserStore struct {
|
||||||
store.UserStore
|
store.UserStore
|
||||||
rootStore *LocalCacheStore
|
rootStore *LocalCacheStore
|
||||||
|
userProfileByIdsMut sync.Mutex
|
||||||
|
userProfileByIdsInvalidations map[string]bool
|
||||||
|
profilesInChannelMut sync.Mutex
|
||||||
|
profilesInChannelInvalidations map[string]bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *LocalCacheUserStore) handleClusterInvalidateScheme(msg *model.ClusterMessage) {
|
func (s *LocalCacheUserStore) handleClusterInvalidateScheme(msg *model.ClusterMessage) {
|
||||||
if msg.Data == ClearCacheMessageData {
|
if msg.Data == ClearCacheMessageData {
|
||||||
s.rootStore.userProfileByIdsCache.Purge()
|
s.rootStore.userProfileByIdsCache.Purge()
|
||||||
} else {
|
} else {
|
||||||
|
s.userProfileByIdsMut.Lock()
|
||||||
|
s.userProfileByIdsInvalidations[msg.Data] = true
|
||||||
|
s.userProfileByIdsMut.Unlock()
|
||||||
s.rootStore.userProfileByIdsCache.Remove(msg.Data)
|
s.rootStore.userProfileByIdsCache.Remove(msg.Data)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -27,11 +37,14 @@ func (s *LocalCacheUserStore) handleClusterInvalidateProfilesInChannel(msg *mode
|
|||||||
if msg.Data == ClearCacheMessageData {
|
if msg.Data == ClearCacheMessageData {
|
||||||
s.rootStore.profilesInChannelCache.Purge()
|
s.rootStore.profilesInChannelCache.Purge()
|
||||||
} else {
|
} else {
|
||||||
|
s.profilesInChannelMut.Lock()
|
||||||
|
s.profilesInChannelInvalidations[msg.Data] = true
|
||||||
|
s.profilesInChannelMut.Unlock()
|
||||||
s.rootStore.profilesInChannelCache.Remove(msg.Data)
|
s.rootStore.profilesInChannelCache.Remove(msg.Data)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) ClearCaches() {
|
func (s *LocalCacheUserStore) ClearCaches() {
|
||||||
s.rootStore.userProfileByIdsCache.Purge()
|
s.rootStore.userProfileByIdsCache.Purge()
|
||||||
s.rootStore.profilesInChannelCache.Purge()
|
s.rootStore.profilesInChannelCache.Purge()
|
||||||
|
|
||||||
@@ -41,7 +54,10 @@ func (s LocalCacheUserStore) ClearCaches() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) InvalidateProfileCacheForUser(userId string) {
|
func (s *LocalCacheUserStore) InvalidateProfileCacheForUser(userId string) {
|
||||||
|
s.userProfileByIdsMut.Lock()
|
||||||
|
s.userProfileByIdsInvalidations[userId] = true
|
||||||
|
s.userProfileByIdsMut.Unlock()
|
||||||
s.rootStore.doInvalidateCacheCluster(s.rootStore.userProfileByIdsCache, userId)
|
s.rootStore.doInvalidateCacheCluster(s.rootStore.userProfileByIdsCache, userId)
|
||||||
|
|
||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
@@ -49,13 +65,16 @@ func (s LocalCacheUserStore) InvalidateProfileCacheForUser(userId string) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) InvalidateProfilesInChannelCacheByUser(userId string) {
|
func (s *LocalCacheUserStore) InvalidateProfilesInChannelCacheByUser(userId string) {
|
||||||
keys, err := s.rootStore.profilesInChannelCache.Keys()
|
keys, err := s.rootStore.profilesInChannelCache.Keys()
|
||||||
if err == nil {
|
if err == nil {
|
||||||
for _, key := range keys {
|
for _, key := range keys {
|
||||||
var userMap map[string]*model.User
|
var userMap map[string]*model.User
|
||||||
if err = s.rootStore.profilesInChannelCache.Get(key, &userMap); err == nil {
|
if err = s.rootStore.profilesInChannelCache.Get(key, &userMap); err == nil {
|
||||||
if _, userInCache := userMap[userId]; userInCache {
|
if _, userInCache := userMap[userId]; userInCache {
|
||||||
|
s.profilesInChannelMut.Lock()
|
||||||
|
s.profilesInChannelInvalidations[key] = true
|
||||||
|
s.profilesInChannelMut.Unlock()
|
||||||
s.rootStore.doInvalidateCacheCluster(s.rootStore.profilesInChannelCache, key)
|
s.rootStore.doInvalidateCacheCluster(s.rootStore.profilesInChannelCache, key)
|
||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
s.rootStore.metrics.IncrementMemCacheInvalidationCounter("Profiles in Channel - Remove by User")
|
s.rootStore.metrics.IncrementMemCacheInvalidationCounter("Profiles in Channel - Remove by User")
|
||||||
@@ -66,14 +85,17 @@ func (s LocalCacheUserStore) InvalidateProfilesInChannelCacheByUser(userId strin
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) InvalidateProfilesInChannelCache(channelId string) {
|
func (s *LocalCacheUserStore) InvalidateProfilesInChannelCache(channelID string) {
|
||||||
s.rootStore.doInvalidateCacheCluster(s.rootStore.profilesInChannelCache, channelId)
|
s.profilesInChannelMut.Lock()
|
||||||
|
s.profilesInChannelInvalidations[channelID] = true
|
||||||
|
s.profilesInChannelMut.Unlock()
|
||||||
|
s.rootStore.doInvalidateCacheCluster(s.rootStore.profilesInChannelCache, channelID)
|
||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
s.rootStore.metrics.IncrementMemCacheInvalidationCounter("Profiles in Channel - Remove by Channel")
|
s.rootStore.metrics.IncrementMemCacheInvalidationCounter("Profiles in Channel - Remove by Channel")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (s *LocalCacheUserStore) GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
if allowFromCache {
|
if allowFromCache {
|
||||||
var cachedMap map[string]*model.User
|
var cachedMap map[string]*model.User
|
||||||
if err := s.rootStore.doStandardReadCache(s.rootStore.profilesInChannelCache, channelId, &cachedMap); err == nil {
|
if err := s.rootStore.doStandardReadCache(s.rootStore.profilesInChannelCache, channelId, &cachedMap); err == nil {
|
||||||
@@ -81,7 +103,16 @@ func (s LocalCacheUserStore) GetAllProfilesInChannel(channelId string, allowFrom
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
userMap, err := s.UserStore.GetAllProfilesInChannel(channelId, allowFromCache)
|
// If it was invalidated, then we need to query master.
|
||||||
|
s.profilesInChannelMut.Lock()
|
||||||
|
if s.profilesInChannelInvalidations[channelId] {
|
||||||
|
ctx = sqlstore.WithMaster(ctx)
|
||||||
|
// And then remove the key from the map.
|
||||||
|
delete(s.profilesInChannelInvalidations, channelId)
|
||||||
|
}
|
||||||
|
s.profilesInChannelMut.Unlock()
|
||||||
|
|
||||||
|
userMap, err := s.UserStore.GetAllProfilesInChannel(ctx, channelId, allowFromCache)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -93,9 +124,9 @@ func (s LocalCacheUserStore) GetAllProfilesInChannel(channelId string, allowFrom
|
|||||||
return userMap, nil
|
return userMap, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s LocalCacheUserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (s *LocalCacheUserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
if !allowFromCache {
|
if !allowFromCache {
|
||||||
return s.UserStore.GetProfileByIds(userIds, options, false)
|
return s.UserStore.GetProfileByIds(ctx, userIds, options, false)
|
||||||
}
|
}
|
||||||
|
|
||||||
if options == nil {
|
if options == nil {
|
||||||
@@ -105,6 +136,7 @@ func (s LocalCacheUserStore) GetProfileByIds(userIds []string, options *store.Us
|
|||||||
users := []*model.User{}
|
users := []*model.User{}
|
||||||
remainingUserIds := make([]string, 0)
|
remainingUserIds := make([]string, 0)
|
||||||
|
|
||||||
|
fromMaster := false
|
||||||
for _, userId := range userIds {
|
for _, userId := range userIds {
|
||||||
var cacheItem *model.User
|
var cacheItem *model.User
|
||||||
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, userId, &cacheItem); err == nil {
|
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, userId, &cacheItem); err == nil {
|
||||||
@@ -112,6 +144,14 @@ func (s LocalCacheUserStore) GetProfileByIds(userIds []string, options *store.Us
|
|||||||
users = append(users, cacheItem)
|
users = append(users, cacheItem)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
|
// If it was invalidated, then we need to query master.
|
||||||
|
s.userProfileByIdsMut.Lock()
|
||||||
|
if s.userProfileByIdsInvalidations[userId] {
|
||||||
|
fromMaster = true
|
||||||
|
// And then remove the key from the map.
|
||||||
|
delete(s.userProfileByIdsInvalidations, userId)
|
||||||
|
}
|
||||||
|
s.userProfileByIdsMut.Unlock()
|
||||||
remainingUserIds = append(remainingUserIds, userId)
|
remainingUserIds = append(remainingUserIds, userId)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -122,7 +162,10 @@ func (s LocalCacheUserStore) GetProfileByIds(userIds []string, options *store.Us
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(remainingUserIds) > 0 {
|
if len(remainingUserIds) > 0 {
|
||||||
remainingUsers, err := s.UserStore.GetProfileByIds(remainingUserIds, options, false)
|
if fromMaster {
|
||||||
|
ctx = sqlstore.WithMaster(ctx)
|
||||||
|
}
|
||||||
|
remainingUsers, err := s.UserStore.GetProfileByIds(ctx, remainingUserIds, options, false)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -139,7 +182,7 @@ func (s LocalCacheUserStore) GetProfileByIds(userIds []string, options *store.Us
|
|||||||
// It checks if the user entry is present in the cache, returning the entry from cache
|
// It checks if the user entry is present in the cache, returning the entry from cache
|
||||||
// if it is present. Otherwise, it fetches the entry from the store and stores it in the
|
// if it is present. Otherwise, it fetches the entry from the store and stores it in the
|
||||||
// cache.
|
// cache.
|
||||||
func (s LocalCacheUserStore) Get(id string) (*model.User, error) {
|
func (s *LocalCacheUserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
var cacheItem *model.User
|
var cacheItem *model.User
|
||||||
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, id, &cacheItem); err == nil {
|
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, id, &cacheItem); err == nil {
|
||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
@@ -150,7 +193,17 @@ func (s LocalCacheUserStore) Get(id string) (*model.User, error) {
|
|||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
s.rootStore.metrics.AddMemCacheMissCounter("Profile By Id", float64(1))
|
s.rootStore.metrics.AddMemCacheMissCounter("Profile By Id", float64(1))
|
||||||
}
|
}
|
||||||
user, err := s.UserStore.Get(id)
|
|
||||||
|
// If it was invalidated, then we need to query master.
|
||||||
|
s.userProfileByIdsMut.Lock()
|
||||||
|
if s.userProfileByIdsInvalidations[id] {
|
||||||
|
ctx = sqlstore.WithMaster(ctx)
|
||||||
|
// And then remove the key from the map.
|
||||||
|
delete(s.userProfileByIdsInvalidations, id)
|
||||||
|
}
|
||||||
|
s.userProfileByIdsMut.Unlock()
|
||||||
|
|
||||||
|
user, err := s.UserStore.Get(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -162,13 +215,14 @@ func (s LocalCacheUserStore) Get(id string) (*model.User, error) {
|
|||||||
// It checks if the user entries are present in the cache, returning the entries from cache
|
// It checks if the user entries are present in the cache, returning the entries from cache
|
||||||
// if it is present. Otherwise, it fetches the entries from the store and stores it in the
|
// if it is present. Otherwise, it fetches the entries from the store and stores it in the
|
||||||
// cache.
|
// cache.
|
||||||
func (s LocalCacheUserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (s *LocalCacheUserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
// we are doing a loop instead of caching the full set in the cache because the number of permutations that we can have
|
// we are doing a loop instead of caching the full set in the cache because the number of permutations that we can have
|
||||||
// in this func is making caching of the total set not beneficial.
|
// in this func is making caching of the total set not beneficial.
|
||||||
var cachedUsers []*model.User
|
var cachedUsers []*model.User
|
||||||
var notCachedUserIds []string
|
var notCachedUserIds []string
|
||||||
uniqIDs := dedup(ids)
|
uniqIDs := dedup(ids)
|
||||||
|
|
||||||
|
fromMaster := false
|
||||||
for _, id := range uniqIDs {
|
for _, id := range uniqIDs {
|
||||||
var cachedUser *model.User
|
var cachedUser *model.User
|
||||||
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, id, &cachedUser); err == nil {
|
if err := s.rootStore.doStandardReadCache(s.rootStore.userProfileByIdsCache, id, &cachedUser); err == nil {
|
||||||
@@ -180,13 +234,24 @@ func (s LocalCacheUserStore) GetMany(ids []string) ([]*model.User, error) {
|
|||||||
if s.rootStore.metrics != nil {
|
if s.rootStore.metrics != nil {
|
||||||
s.rootStore.metrics.AddMemCacheMissCounter("Profile By Id", float64(1))
|
s.rootStore.metrics.AddMemCacheMissCounter("Profile By Id", float64(1))
|
||||||
}
|
}
|
||||||
|
// If it was invalidated, then we need to query master.
|
||||||
|
s.userProfileByIdsMut.Lock()
|
||||||
|
if s.userProfileByIdsInvalidations[id] {
|
||||||
|
fromMaster = true
|
||||||
|
// And then remove the key from the map.
|
||||||
|
delete(s.userProfileByIdsInvalidations, id)
|
||||||
|
}
|
||||||
|
s.userProfileByIdsMut.Unlock()
|
||||||
|
|
||||||
notCachedUserIds = append(notCachedUserIds, id)
|
notCachedUserIds = append(notCachedUserIds, id)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(notCachedUserIds) > 0 {
|
if len(notCachedUserIds) > 0 {
|
||||||
dbUsers, err := s.UserStore.GetMany(notCachedUserIds)
|
if fromMaster {
|
||||||
|
ctx = sqlstore.WithMaster(ctx)
|
||||||
|
}
|
||||||
|
dbUsers, err := s.UserStore.GetMany(ctx, notCachedUserIds)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,12 +4,14 @@
|
|||||||
package localcachelayer
|
package localcachelayer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
"github.com/mattermost/mattermost-server/v5/model"
|
"github.com/mattermost/mattermost-server/v5/model"
|
||||||
|
"github.com/mattermost/mattermost-server/v5/plugin/plugintest/mock"
|
||||||
"github.com/mattermost/mattermost-server/v5/store"
|
"github.com/mattermost/mattermost-server/v5/store"
|
||||||
"github.com/mattermost/mattermost-server/v5/store/storetest"
|
"github.com/mattermost/mattermost-server/v5/store/storetest"
|
||||||
"github.com/mattermost/mattermost-server/v5/store/storetest/mocks"
|
"github.com/mattermost/mattermost-server/v5/store/storetest/mocks"
|
||||||
@@ -33,12 +35,12 @@ func TestUserStoreCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUser, err := cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
gotUser, err := cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeUser, gotUser)
|
assert.Equal(t, fakeUser, gotUser)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
_, _ = cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -48,12 +50,12 @@ func TestUserStoreCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUser, err := cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
gotUser, err := cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeUser, gotUser)
|
assert.Equal(t, fakeUser, gotUser)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 1)
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, false)
|
_, _ = cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, false)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 2)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -63,13 +65,13 @@ func TestUserStoreCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUser, err := cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
gotUser, err := cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeUser, gotUser)
|
assert.Equal(t, fakeUser, gotUser)
|
||||||
|
|
||||||
cachedStore.User().InvalidateProfileCacheForUser("123")
|
cachedStore.User().InvalidateProfileCacheForUser("123")
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
_, _ = cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetProfileByIds", 2)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -79,7 +81,7 @@ func TestUserStoreCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
storedUsers, err := mockStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, false)
|
storedUsers, err := mockStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, false)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
originalProps := make([]model.StringMap, len(storedUsers))
|
originalProps := make([]model.StringMap, len(storedUsers))
|
||||||
@@ -90,14 +92,14 @@ func TestUserStoreCache(t *testing.T) {
|
|||||||
storedUsers[i].NotifyProps["key"] = "somevalue"
|
storedUsers[i].NotifyProps["key"] = "somevalue"
|
||||||
}
|
}
|
||||||
|
|
||||||
cachedUsers, err := cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
cachedUsers, err := cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
for i := 0; i < len(storedUsers); i++ {
|
for i := 0; i < len(storedUsers); i++ {
|
||||||
assert.Equal(t, storedUsers[i].Id, cachedUsers[i].Id)
|
assert.Equal(t, storedUsers[i].Id, cachedUsers[i].Id)
|
||||||
}
|
}
|
||||||
|
|
||||||
cachedUsers, err = cachedStore.User().GetProfileByIds(fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
cachedUsers, err = cachedStore.User().GetProfileByIds(context.Background(), fakeUserIds, &store.UserGetByIdsOpts{}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
for i := 0; i < len(storedUsers); i++ {
|
for i := 0; i < len(storedUsers); i++ {
|
||||||
storedUsers[i].Props = model.StringMap{}
|
storedUsers[i].Props = model.StringMap{}
|
||||||
@@ -129,12 +131,12 @@ func TestUserStoreProfilesInChannelCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotMap, err := cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
gotMap, err := cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeMap, gotMap)
|
assert.Equal(t, fakeMap, gotMap)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
_, _ = cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -144,12 +146,12 @@ func TestUserStoreProfilesInChannelCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotMap, err := cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
gotMap, err := cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeMap, gotMap)
|
assert.Equal(t, fakeMap, gotMap)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetAllProfilesInChannel(fakeChannelId, false)
|
_, _ = cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, false)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -159,14 +161,14 @@ func TestUserStoreProfilesInChannelCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotMap, err := cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
gotMap, err := cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeMap, gotMap)
|
assert.Equal(t, fakeMap, gotMap)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
||||||
|
|
||||||
cachedStore.User().InvalidateProfilesInChannelCache("123")
|
cachedStore.User().InvalidateProfilesInChannelCache("123")
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
_, _ = cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -176,14 +178,14 @@ func TestUserStoreProfilesInChannelCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotMap, err := cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
gotMap, err := cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeMap, gotMap)
|
assert.Equal(t, fakeMap, gotMap)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 1)
|
||||||
|
|
||||||
cachedStore.User().InvalidateProfilesInChannelCacheByUser("456")
|
cachedStore.User().InvalidateProfilesInChannelCacheByUser("456")
|
||||||
|
|
||||||
_, _ = cachedStore.User().GetAllProfilesInChannel(fakeChannelId, true)
|
_, _ = cachedStore.User().GetAllProfilesInChannel(context.Background(), fakeChannelId, true)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetAllProfilesInChannel", 2)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -201,12 +203,12 @@ func TestUserStoreGetCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUser, err := cachedStore.User().Get(fakeUserId)
|
gotUser, err := cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeUser, gotUser)
|
assert.Equal(t, fakeUser, gotUser)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
||||||
|
|
||||||
_, _ = cachedStore.User().Get(fakeUserId)
|
_, _ = cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -216,14 +218,14 @@ func TestUserStoreGetCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUser, err := cachedStore.User().Get(fakeUserId)
|
gotUser, err := cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, fakeUser, gotUser)
|
assert.Equal(t, fakeUser, gotUser)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 1)
|
||||||
|
|
||||||
cachedStore.User().InvalidateProfileCacheForUser("123")
|
cachedStore.User().InvalidateProfileCacheForUser("123")
|
||||||
|
|
||||||
_, _ = cachedStore.User().Get(fakeUserId)
|
_, _ = cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "Get", 2)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -233,20 +235,20 @@ func TestUserStoreGetCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
storedUser, err := mockStore.User().Get(fakeUserId)
|
storedUser, err := mockStore.User().Get(context.Background(), fakeUserId)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
originalProps := storedUser.NotifyProps
|
originalProps := storedUser.NotifyProps
|
||||||
|
|
||||||
storedUser.NotifyProps = map[string]string{}
|
storedUser.NotifyProps = map[string]string{}
|
||||||
storedUser.NotifyProps["key"] = "somevalue"
|
storedUser.NotifyProps["key"] = "somevalue"
|
||||||
|
|
||||||
cachedUser, err := cachedStore.User().Get(fakeUserId)
|
cachedUser, err := cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, storedUser, cachedUser)
|
assert.Equal(t, storedUser, cachedUser)
|
||||||
|
|
||||||
storedUser.Props = model.StringMap{}
|
storedUser.Props = model.StringMap{}
|
||||||
storedUser.Timezone = model.StringMap{}
|
storedUser.Timezone = model.StringMap{}
|
||||||
cachedUser, err = cachedStore.User().Get(fakeUserId)
|
cachedUser, err = cachedStore.User().Get(context.Background(), fakeUserId)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, storedUser, cachedUser)
|
assert.Equal(t, storedUser, cachedUser)
|
||||||
if storedUser == cachedUser {
|
if storedUser == cachedUser {
|
||||||
@@ -276,13 +278,13 @@ func TestUserStoreGetManyCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUsers, err := cachedStore.User().GetMany([]string{fakeUser.Id, otherFakeUser.Id})
|
gotUsers, err := cachedStore.User().GetMany(context.Background(), []string{fakeUser.Id, otherFakeUser.Id})
|
||||||
require.Nil(t, err)
|
require.Nil(t, err)
|
||||||
assert.Len(t, gotUsers, 2)
|
assert.Len(t, gotUsers, 2)
|
||||||
assert.Contains(t, gotUsers, fakeUser)
|
assert.Contains(t, gotUsers, fakeUser)
|
||||||
assert.Contains(t, gotUsers, otherFakeUser)
|
assert.Contains(t, gotUsers, otherFakeUser)
|
||||||
|
|
||||||
gotUsers, err = cachedStore.User().GetMany([]string{fakeUser.Id, otherFakeUser.Id})
|
gotUsers, err = cachedStore.User().GetMany(context.Background(), []string{fakeUser.Id, otherFakeUser.Id})
|
||||||
require.Nil(t, err)
|
require.Nil(t, err)
|
||||||
assert.Len(t, gotUsers, 2)
|
assert.Len(t, gotUsers, 2)
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetMany", 1)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetMany", 1)
|
||||||
@@ -294,7 +296,7 @@ func TestUserStoreGetManyCache(t *testing.T) {
|
|||||||
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
cachedStore, err := NewLocalCacheLayer(mockStore, nil, nil, mockCacheProvider)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
gotUsers, err := cachedStore.User().GetMany([]string{fakeUser.Id, otherFakeUser.Id})
|
gotUsers, err := cachedStore.User().GetMany(context.Background(), []string{fakeUser.Id, otherFakeUser.Id})
|
||||||
require.Nil(t, err)
|
require.Nil(t, err)
|
||||||
assert.Len(t, gotUsers, 2)
|
assert.Len(t, gotUsers, 2)
|
||||||
assert.Contains(t, gotUsers, fakeUser)
|
assert.Contains(t, gotUsers, fakeUser)
|
||||||
@@ -302,10 +304,10 @@ func TestUserStoreGetManyCache(t *testing.T) {
|
|||||||
|
|
||||||
cachedStore.User().InvalidateProfileCacheForUser("123")
|
cachedStore.User().InvalidateProfileCacheForUser("123")
|
||||||
|
|
||||||
gotUsers, err = cachedStore.User().GetMany([]string{fakeUser.Id, otherFakeUser.Id})
|
gotUsers, err = cachedStore.User().GetMany(context.Background(), []string{fakeUser.Id, otherFakeUser.Id})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Len(t, gotUsers, 2)
|
assert.Len(t, gotUsers, 2)
|
||||||
mockStore.User().(*mocks.UserStore).AssertCalled(t, "GetMany", []string{"123"})
|
mockStore.User().(*mocks.UserStore).AssertCalled(t, "GetMany", mock.Anything, []string{"123"})
|
||||||
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetMany", 2)
|
mockStore.User().(*mocks.UserStore).AssertNumberOfCalls(t, "GetMany", 2)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8394,7 +8394,7 @@ func (s *OpenTracingLayerUserStore) DeactivateGuests() ([]string, error) {
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *OpenTracingLayerUserStore) DemoteUserToGuest(userID string) error {
|
func (s *OpenTracingLayerUserStore) DemoteUserToGuest(userID string) (*model.User, error) {
|
||||||
origCtx := s.Root.Store.Context()
|
origCtx := s.Root.Store.Context()
|
||||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.DemoteUserToGuest")
|
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.DemoteUserToGuest")
|
||||||
s.Root.Store.SetContext(newCtx)
|
s.Root.Store.SetContext(newCtx)
|
||||||
@@ -8403,16 +8403,16 @@ func (s *OpenTracingLayerUserStore) DemoteUserToGuest(userID string) error {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
defer span.Finish()
|
defer span.Finish()
|
||||||
err := s.UserStore.DemoteUserToGuest(userID)
|
result, err := s.UserStore.DemoteUserToGuest(userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
span.LogFields(spanlog.Error(err))
|
span.LogFields(spanlog.Error(err))
|
||||||
ext.Error.Set(span, true)
|
ext.Error.Set(span, true)
|
||||||
}
|
}
|
||||||
|
|
||||||
return err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *OpenTracingLayerUserStore) Get(id string) (*model.User, error) {
|
func (s *OpenTracingLayerUserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
origCtx := s.Root.Store.Context()
|
origCtx := s.Root.Store.Context()
|
||||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.Get")
|
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.Get")
|
||||||
s.Root.Store.SetContext(newCtx)
|
s.Root.Store.SetContext(newCtx)
|
||||||
@@ -8421,7 +8421,7 @@ func (s *OpenTracingLayerUserStore) Get(id string) (*model.User, error) {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
defer span.Finish()
|
defer span.Finish()
|
||||||
result, err := s.UserStore.Get(id)
|
result, err := s.UserStore.Get(ctx, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
span.LogFields(spanlog.Error(err))
|
span.LogFields(spanlog.Error(err))
|
||||||
ext.Error.Set(span, true)
|
ext.Error.Set(span, true)
|
||||||
@@ -8502,7 +8502,7 @@ func (s *OpenTracingLayerUserStore) GetAllProfiles(options *model.UserGetOptions
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *OpenTracingLayerUserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (s *OpenTracingLayerUserStore) GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
origCtx := s.Root.Store.Context()
|
origCtx := s.Root.Store.Context()
|
||||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetAllProfilesInChannel")
|
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetAllProfilesInChannel")
|
||||||
s.Root.Store.SetContext(newCtx)
|
s.Root.Store.SetContext(newCtx)
|
||||||
@@ -8511,7 +8511,7 @@ func (s *OpenTracingLayerUserStore) GetAllProfilesInChannel(channelId string, al
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
defer span.Finish()
|
defer span.Finish()
|
||||||
result, err := s.UserStore.GetAllProfilesInChannel(channelId, allowFromCache)
|
result, err := s.UserStore.GetAllProfilesInChannel(ctx, channelId, allowFromCache)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
span.LogFields(spanlog.Error(err))
|
span.LogFields(spanlog.Error(err))
|
||||||
ext.Error.Set(span, true)
|
ext.Error.Set(span, true)
|
||||||
@@ -8703,7 +8703,7 @@ func (s *OpenTracingLayerUserStore) GetKnownUsers(userID string) ([]string, erro
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *OpenTracingLayerUserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (s *OpenTracingLayerUserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
origCtx := s.Root.Store.Context()
|
origCtx := s.Root.Store.Context()
|
||||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetMany")
|
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetMany")
|
||||||
s.Root.Store.SetContext(newCtx)
|
s.Root.Store.SetContext(newCtx)
|
||||||
@@ -8712,7 +8712,7 @@ func (s *OpenTracingLayerUserStore) GetMany(ids []string) ([]*model.User, error)
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
defer span.Finish()
|
defer span.Finish()
|
||||||
result, err := s.UserStore.GetMany(ids)
|
result, err := s.UserStore.GetMany(ctx, ids)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
span.LogFields(spanlog.Error(err))
|
span.LogFields(spanlog.Error(err))
|
||||||
ext.Error.Set(span, true)
|
ext.Error.Set(span, true)
|
||||||
@@ -8757,7 +8757,7 @@ func (s *OpenTracingLayerUserStore) GetProfileByGroupChannelIdsForUser(userId st
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *OpenTracingLayerUserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (s *OpenTracingLayerUserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
origCtx := s.Root.Store.Context()
|
origCtx := s.Root.Store.Context()
|
||||||
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetProfileByIds")
|
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UserStore.GetProfileByIds")
|
||||||
s.Root.Store.SetContext(newCtx)
|
s.Root.Store.SetContext(newCtx)
|
||||||
@@ -8766,7 +8766,7 @@ func (s *OpenTracingLayerUserStore) GetProfileByIds(userIds []string, options *s
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
defer span.Finish()
|
defer span.Finish()
|
||||||
result, err := s.UserStore.GetProfileByIds(userIds, options, allowFromCache)
|
result, err := s.UserStore.GetProfileByIds(ctx, userIds, options, allowFromCache)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
span.LogFields(spanlog.Error(err))
|
span.LogFields(spanlog.Error(err))
|
||||||
ext.Error.Set(span, true)
|
ext.Error.Set(span, true)
|
||||||
|
|||||||
@@ -9110,31 +9110,31 @@ func (s *RetryLayerUserStore) DeactivateGuests() ([]string, error) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RetryLayerUserStore) DemoteUserToGuest(userID string) error {
|
func (s *RetryLayerUserStore) DemoteUserToGuest(userID string) (*model.User, error) {
|
||||||
|
|
||||||
tries := 0
|
tries := 0
|
||||||
for {
|
for {
|
||||||
err := s.UserStore.DemoteUserToGuest(userID)
|
result, err := s.UserStore.DemoteUserToGuest(userID)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return nil
|
return result, nil
|
||||||
}
|
}
|
||||||
if !isRepeatableError(err) {
|
if !isRepeatableError(err) {
|
||||||
return err
|
return result, err
|
||||||
}
|
}
|
||||||
tries++
|
tries++
|
||||||
if tries >= 3 {
|
if tries >= 3 {
|
||||||
err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures")
|
err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures")
|
||||||
return err
|
return result, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RetryLayerUserStore) Get(id string) (*model.User, error) {
|
func (s *RetryLayerUserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
|
|
||||||
tries := 0
|
tries := 0
|
||||||
for {
|
for {
|
||||||
result, err := s.UserStore.Get(id)
|
result, err := s.UserStore.Get(ctx, id)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
@@ -9230,11 +9230,11 @@ func (s *RetryLayerUserStore) GetAllProfiles(options *model.UserGetOptions) ([]*
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RetryLayerUserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (s *RetryLayerUserStore) GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
|
|
||||||
tries := 0
|
tries := 0
|
||||||
for {
|
for {
|
||||||
result, err := s.UserStore.GetAllProfilesInChannel(channelId, allowFromCache)
|
result, err := s.UserStore.GetAllProfilesInChannel(ctx, channelId, allowFromCache)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
@@ -9428,11 +9428,11 @@ func (s *RetryLayerUserStore) GetKnownUsers(userID string) ([]string, error) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RetryLayerUserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (s *RetryLayerUserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
|
|
||||||
tries := 0
|
tries := 0
|
||||||
for {
|
for {
|
||||||
result, err := s.UserStore.GetMany(ids)
|
result, err := s.UserStore.GetMany(ctx, ids)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
@@ -9488,11 +9488,11 @@ func (s *RetryLayerUserStore) GetProfileByGroupChannelIdsForUser(userId string,
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *RetryLayerUserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (s *RetryLayerUserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
|
|
||||||
tries := 0
|
tries := 0
|
||||||
for {
|
for {
|
||||||
result, err := s.UserStore.GetProfileByIds(userIds, options, allowFromCache)
|
result, err := s.UserStore.GetProfileByIds(ctx, userIds, options, allowFromCache)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,8 @@
|
|||||||
package searchlayer
|
package searchlayer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
"github.com/pkg/errors"
|
"github.com/pkg/errors"
|
||||||
|
|
||||||
"github.com/mattermost/mattermost-server/v5/mlog"
|
"github.com/mattermost/mattermost-server/v5/mlog"
|
||||||
@@ -193,7 +195,7 @@ func (c *SearchChannelStore) PermanentDeleteMembersByUser(userId string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *SearchChannelStore) RemoveAllDeactivatedMembers(channelId string) error {
|
func (c *SearchChannelStore) RemoveAllDeactivatedMembers(channelId string) error {
|
||||||
profiles, errProfiles := c.rootStore.User().GetAllProfilesInChannel(channelId, true)
|
profiles, errProfiles := c.rootStore.User().GetAllProfilesInChannel(context.Background(), channelId, true)
|
||||||
if errProfiles != nil {
|
if errProfiles != nil {
|
||||||
mlog.Warn("Encountered error indexing users for channel", mlog.String("channel_id", channelId), mlog.Err(errProfiles))
|
mlog.Warn("Encountered error indexing users for channel", mlog.String("channel_id", channelId), mlog.Err(errProfiles))
|
||||||
}
|
}
|
||||||
@@ -210,7 +212,7 @@ func (c *SearchChannelStore) RemoveAllDeactivatedMembers(channelId string) error
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *SearchChannelStore) PermanentDeleteMembersByChannel(channelId string) error {
|
func (c *SearchChannelStore) PermanentDeleteMembersByChannel(channelId string) error {
|
||||||
profiles, errProfiles := c.rootStore.User().GetAllProfilesInChannel(channelId, true)
|
profiles, errProfiles := c.rootStore.User().GetAllProfilesInChannel(context.Background(), channelId, true)
|
||||||
if errProfiles != nil {
|
if errProfiles != nil {
|
||||||
mlog.Warn("Encountered error indexing users for channel", mlog.String("channel_id", channelId), mlog.Err(errProfiles))
|
mlog.Warn("Encountered error indexing users for channel", mlog.String("channel_id", channelId), mlog.Err(errProfiles))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package searchlayer
|
package searchlayer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/mattermost/mattermost-server/v5/mlog"
|
"github.com/mattermost/mattermost-server/v5/mlog"
|
||||||
@@ -67,7 +68,7 @@ func (s *SearchStore) User() store.UserStore {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *SearchStore) indexUserFromID(userId string) {
|
func (s *SearchStore) indexUserFromID(userId string) {
|
||||||
user, err := s.User().Get(userId)
|
user, err := s.User().Get(context.Background(), userId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package searchlayer
|
package searchlayer
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/pkg/errors"
|
"github.com/pkg/errors"
|
||||||
@@ -54,7 +55,7 @@ func (s *SearchUserStore) Search(teamId, term string, options *model.UserSearchO
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
users, nErr := s.UserStore.GetProfileByIds(usersIds, nil, false)
|
users, nErr := s.UserStore.GetProfileByIds(context.Background(), usersIds, nil, false)
|
||||||
if nErr != nil {
|
if nErr != nil {
|
||||||
mlog.Warn("Encountered error on Search", mlog.String("search_engine", engine.GetName()), mlog.Err(nErr))
|
mlog.Warn("Encountered error on Search", mlog.String("search_engine", engine.GetName()), mlog.Err(nErr))
|
||||||
continue
|
continue
|
||||||
@@ -89,7 +90,7 @@ func (s *SearchUserStore) Save(user *model.User) (*model.User, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *SearchUserStore) PermanentDelete(userId string) error {
|
func (s *SearchUserStore) PermanentDelete(userId string) error {
|
||||||
user, userErr := s.UserStore.Get(userId)
|
user, userErr := s.UserStore.Get(context.Background(), userId)
|
||||||
if userErr != nil {
|
if userErr != nil {
|
||||||
mlog.Warn("Encountered error deleting user", mlog.String("user_id", userId), mlog.Err(userErr))
|
mlog.Warn("Encountered error deleting user", mlog.String("user_id", userId), mlog.Err(userErr))
|
||||||
}
|
}
|
||||||
@@ -116,14 +117,14 @@ func (s *SearchUserStore) autocompleteUsersInChannelByEngine(engine searchengine
|
|||||||
|
|
||||||
uchan := make(chan store.StoreResult, 1)
|
uchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
users, nErr := s.UserStore.GetProfileByIds(uchanIds, nil, false)
|
users, nErr := s.UserStore.GetProfileByIds(context.Background(), uchanIds, nil, false)
|
||||||
uchan <- store.StoreResult{Data: users, NErr: nErr}
|
uchan <- store.StoreResult{Data: users, NErr: nErr}
|
||||||
close(uchan)
|
close(uchan)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
nuchan := make(chan store.StoreResult, 1)
|
nuchan := make(chan store.StoreResult, 1)
|
||||||
go func() {
|
go func() {
|
||||||
users, nErr := s.UserStore.GetProfileByIds(nuchanIds, nil, false)
|
users, nErr := s.UserStore.GetProfileByIds(context.Background(), nuchanIds, nil, false)
|
||||||
nuchan <- store.StoreResult{Data: users, NErr: nErr}
|
nuchan <- store.StoreResult{Data: users, NErr: nErr}
|
||||||
close(nuchan)
|
close(nuchan)
|
||||||
}()
|
}()
|
||||||
|
|||||||
@@ -18,8 +18,8 @@ const (
|
|||||||
useMaster contextValue = "useMaster"
|
useMaster contextValue = "useMaster"
|
||||||
)
|
)
|
||||||
|
|
||||||
// withMaster adds the context value that master DB should be selected for this request.
|
// WithMaster adds the context value that master DB should be selected for this request.
|
||||||
func withMaster(ctx context.Context) context.Context {
|
func WithMaster(ctx context.Context) context.Context {
|
||||||
return context.WithValue(ctx, storeContextKey(useMaster), true)
|
return context.WithValue(ctx, storeContextKey(useMaster), true)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,6 @@ import (
|
|||||||
func TestContextMaster(t *testing.T) {
|
func TestContextMaster(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
m := withMaster(ctx)
|
m := WithMaster(ctx)
|
||||||
assert.True(t, hasMaster(m))
|
assert.True(t, hasMaster(m))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package sqlstore
|
package sqlstore
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"regexp"
|
"regexp"
|
||||||
@@ -677,7 +678,7 @@ func (s *SqlPostStore) prepareThreadedResponse(posts []*postWithExtra, extended,
|
|||||||
var users []*model.User
|
var users []*model.User
|
||||||
if extended {
|
if extended {
|
||||||
var err error
|
var err error
|
||||||
users, err = s.User().GetProfileByIds(userIds, &store.UserGetByIdsOpts{}, true)
|
users, err = s.User().GetProfileByIds(context.Background(), userIds, &store.UserGetByIdsOpts{}, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ func (me SqlSessionStore) Get(sessionIdOrToken string) (*model.Session, error) {
|
|||||||
session := sessions[0]
|
session := sessions[0]
|
||||||
|
|
||||||
tempMembers, err := me.Team().GetTeamsForUser(
|
tempMembers, err := me.Team().GetTeamsForUser(
|
||||||
withMaster(context.Background()),
|
WithMaster(context.Background()),
|
||||||
session.UserId)
|
session.UserId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
|
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package sqlstore
|
package sqlstore
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -288,7 +289,7 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get
|
|||||||
var users []*model.User
|
var users []*model.User
|
||||||
if opts.Extended {
|
if opts.Extended {
|
||||||
var err error
|
var err error
|
||||||
users, err = s.User().GetProfileByIds(userIds, &store.UserGetByIdsOpts{}, true)
|
users, err = s.User().GetProfileByIds(context.Background(), userIds, &store.UserGetByIdsOpts{}, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrapf(err, "failed to get threads for user id=%s", userId)
|
return nil, errors.Wrapf(err, "failed to get threads for user id=%s", userId)
|
||||||
}
|
}
|
||||||
@@ -376,7 +377,7 @@ func (s *SqlThreadStore) GetThreadForUser(userId, teamId, threadId string, exten
|
|||||||
var users []*model.User
|
var users []*model.User
|
||||||
if extended {
|
if extended {
|
||||||
var err error
|
var err error
|
||||||
users, err = s.User().GetProfileByIds(thread.Participants, &store.UserGetByIdsOpts{}, true)
|
users, err = s.User().GetProfileByIds(context.Background(), thread.Participants, &store.UserGetByIdsOpts{}, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrapf(err, "failed to get threads for user id=%s", userId)
|
return nil, errors.Wrapf(err, "failed to get threads for user id=%s", userId)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package sqlstore
|
package sqlstore
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -38,7 +39,7 @@ type SqlUserStore struct {
|
|||||||
usersQuery sq.SelectBuilder
|
usersQuery sq.SelectBuilder
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) ClearCaches() {}
|
func (us *SqlUserStore) ClearCaches() {}
|
||||||
|
|
||||||
func (us SqlUserStore) InvalidateProfileCacheForUser(userId string) {}
|
func (us SqlUserStore) InvalidateProfileCacheForUser(userId string) {}
|
||||||
|
|
||||||
@@ -326,28 +327,41 @@ func (us SqlUserStore) UpdateMfaActive(userId string, active bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetMany returns a list of users for the provided list of ids
|
// GetMany returns a list of users for the provided list of ids
|
||||||
func (us SqlUserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (us SqlUserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
query := us.usersQuery.Where(sq.Eq{"Id": ids})
|
query := us.usersQuery.Where(sq.Eq{"Id": ids})
|
||||||
queryString, args, err := query.ToSql()
|
queryString, args, err := query.ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrap(err, "users_get_many_tosql")
|
return nil, errors.Wrap(err, "users_get_many_tosql")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var db *gorp.DbMap
|
||||||
|
if hasMaster(ctx) {
|
||||||
|
db = us.GetMaster()
|
||||||
|
} else {
|
||||||
|
db = us.GetReplica()
|
||||||
|
}
|
||||||
|
|
||||||
var users []*model.User
|
var users []*model.User
|
||||||
if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil {
|
if _, err := db.Select(&users, queryString, args...); err != nil {
|
||||||
return nil, errors.Wrap(err, "users_get_many_select")
|
return nil, errors.Wrap(err, "users_get_many_select")
|
||||||
}
|
}
|
||||||
|
|
||||||
return users, nil
|
return users, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) Get(id string) (*model.User, error) {
|
func (us SqlUserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
query := us.usersQuery.Where("Id = ?", id)
|
query := us.usersQuery.Where("Id = ?", id)
|
||||||
queryString, args, err := query.ToSql()
|
queryString, args, err := query.ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrap(err, "users_get_tosql")
|
return nil, errors.Wrap(err, "users_get_tosql")
|
||||||
}
|
}
|
||||||
row := us.GetReplica().Db.QueryRow(queryString, args...)
|
var db *gorp.DbMap
|
||||||
|
if hasMaster(ctx) {
|
||||||
|
db = us.GetMaster()
|
||||||
|
} else {
|
||||||
|
db = us.GetReplica()
|
||||||
|
}
|
||||||
|
row := db.Db.QueryRow(queryString, args...)
|
||||||
|
|
||||||
var user model.User
|
var user model.User
|
||||||
var props, notifyProps, timezone []byte
|
var props, notifyProps, timezone []byte
|
||||||
@@ -703,10 +717,10 @@ func (us SqlUserStore) GetProfilesInChannelByStatus(options *model.UserGetOption
|
|||||||
return users, nil
|
return users, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (us SqlUserStore) GetAllProfilesInChannel(ctx context.Context, channelID string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
query := us.usersQuery.
|
query := us.usersQuery.
|
||||||
Join("ChannelMembers cm ON ( cm.UserId = u.Id )").
|
Join("ChannelMembers cm ON ( cm.UserId = u.Id )").
|
||||||
Where("cm.ChannelId = ?", channelId).
|
Where("cm.ChannelId = ?", channelID).
|
||||||
Where("u.DeleteAt = 0").
|
Where("u.DeleteAt = 0").
|
||||||
OrderBy("u.Username ASC")
|
OrderBy("u.Username ASC")
|
||||||
|
|
||||||
@@ -714,8 +728,15 @@ func (us SqlUserStore) GetAllProfilesInChannel(channelId string, allowFromCache
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrap(err, "get_all_profiles_in_channel_tosql")
|
return nil, errors.Wrap(err, "get_all_profiles_in_channel_tosql")
|
||||||
}
|
}
|
||||||
|
var db *gorp.DbMap
|
||||||
|
if hasMaster(ctx) {
|
||||||
|
db = us.GetMaster()
|
||||||
|
} else {
|
||||||
|
db = us.GetReplica()
|
||||||
|
}
|
||||||
|
|
||||||
var users []*model.User
|
var users []*model.User
|
||||||
rows, err := us.GetReplica().Db.Query(queryString, args...)
|
rows, err := db.Db.Query(queryString, args...)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.Wrap(err, "failed to find Users")
|
return nil, errors.Wrap(err, "failed to find Users")
|
||||||
}
|
}
|
||||||
@@ -914,7 +935,7 @@ func (us SqlUserStore) GetNewUsersForTeam(teamId string, offset, limit int, view
|
|||||||
return users, nil
|
return users, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (us SqlUserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
if options == nil {
|
if options == nil {
|
||||||
options = &store.UserGetByIdsOpts{}
|
options = &store.UserGetByIdsOpts{}
|
||||||
}
|
}
|
||||||
@@ -939,7 +960,14 @@ func (us SqlUserStore) GetProfileByIds(userIds []string, options *store.UserGetB
|
|||||||
return nil, errors.Wrap(err, "get_profile_by_ids_tosql")
|
return nil, errors.Wrap(err, "get_profile_by_ids_tosql")
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := us.GetReplica().Select(&users, queryString, args...); err != nil {
|
var db *gorp.DbMap
|
||||||
|
if hasMaster(ctx) {
|
||||||
|
db = us.GetMaster()
|
||||||
|
} else {
|
||||||
|
db = us.GetReplica()
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := db.Select(&users, queryString, args...); err != nil {
|
||||||
return nil, errors.Wrap(err, "failed to find Users")
|
return nil, errors.Wrap(err, "failed to find Users")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1775,7 +1803,7 @@ func (us SqlUserStore) PromoteGuestToUser(userId string) error {
|
|||||||
}
|
}
|
||||||
defer finalizeTransaction(transaction)
|
defer finalizeTransaction(transaction)
|
||||||
|
|
||||||
user, err := us.Get(userId)
|
user, err := us.Get(context.Background(), userId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -1837,76 +1865,80 @@ func (us SqlUserStore) PromoteGuestToUser(userId string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) DemoteUserToGuest(userId string) error {
|
func (us SqlUserStore) DemoteUserToGuest(userID string) (*model.User, error) {
|
||||||
transaction, err := us.GetMaster().Begin()
|
transaction, err := us.GetMaster().Begin()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.Wrap(err, "begin_transaction")
|
return nil, errors.Wrap(err, "begin_transaction")
|
||||||
}
|
}
|
||||||
defer finalizeTransaction(transaction)
|
defer finalizeTransaction(transaction)
|
||||||
|
|
||||||
user, err := us.Get(userId)
|
user, err := us.Get(context.Background(), userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
roles := user.GetRoles()
|
roles := user.GetRoles()
|
||||||
|
|
||||||
newRoles := []string{}
|
newRoles := []string{}
|
||||||
for _, role := range roles {
|
for _, role := range roles {
|
||||||
if role == "system_user" {
|
if role == model.SYSTEM_USER_ROLE_ID {
|
||||||
newRoles = append(newRoles, "system_guest")
|
newRoles = append(newRoles, model.SYSTEM_GUEST_ROLE_ID)
|
||||||
} else if role != "system_admin" {
|
} else if role != model.SYSTEM_ADMIN_ROLE_ID {
|
||||||
newRoles = append(newRoles, role)
|
newRoles = append(newRoles, role)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
curTime := model.GetMillis()
|
curTime := model.GetMillis()
|
||||||
|
newRolesDBStr := strings.Join(newRoles, " ")
|
||||||
query := us.getQueryBuilder().Update("Users").
|
query := us.getQueryBuilder().Update("Users").
|
||||||
Set("Roles", strings.Join(newRoles, " ")).
|
Set("Roles", newRolesDBStr).
|
||||||
Set("UpdateAt", curTime).
|
Set("UpdateAt", curTime).
|
||||||
Where(sq.Eq{"Id": userId})
|
Where(sq.Eq{"Id": userID})
|
||||||
|
|
||||||
queryString, args, err := query.ToSql()
|
queryString, args, err := query.ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.Wrap(err, "demote_user_to_guest_tosql")
|
return nil, errors.Wrap(err, "demote_user_to_guest_tosql")
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err = transaction.Exec(queryString, args...); err != nil {
|
if _, err = transaction.Exec(queryString, args...); err != nil {
|
||||||
return errors.Wrapf(err, "failed to update User with userId=%s", userId)
|
return nil, errors.Wrapf(err, "failed to update User with userId=%s", userID)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
user.Roles = newRolesDBStr
|
||||||
|
user.UpdateAt = curTime
|
||||||
|
|
||||||
query = us.getQueryBuilder().Update("ChannelMembers").
|
query = us.getQueryBuilder().Update("ChannelMembers").
|
||||||
Set("SchemeUser", false).
|
Set("SchemeUser", false).
|
||||||
Set("SchemeGuest", true).
|
Set("SchemeGuest", true).
|
||||||
Where(sq.Eq{"UserId": userId})
|
Where(sq.Eq{"UserId": userID})
|
||||||
|
|
||||||
queryString, args, err = query.ToSql()
|
queryString, args, err = query.ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.Wrap(err, "demote_user_to_guest_tosql")
|
return nil, errors.Wrap(err, "demote_user_to_guest_tosql")
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err = transaction.Exec(queryString, args...); err != nil {
|
if _, err = transaction.Exec(queryString, args...); err != nil {
|
||||||
return errors.Wrapf(err, "failed to update ChannelMembers with userId=%s", userId)
|
return nil, errors.Wrapf(err, "failed to update ChannelMembers with userId=%s", userID)
|
||||||
}
|
}
|
||||||
|
|
||||||
query = us.getQueryBuilder().Update("TeamMembers").
|
query = us.getQueryBuilder().Update("TeamMembers").
|
||||||
Set("SchemeUser", false).
|
Set("SchemeUser", false).
|
||||||
Set("SchemeGuest", true).
|
Set("SchemeGuest", true).
|
||||||
Where(sq.Eq{"UserId": userId})
|
Where(sq.Eq{"UserId": userID})
|
||||||
|
|
||||||
queryString, args, err = query.ToSql()
|
queryString, args, err = query.ToSql()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.Wrap(err, "demote_user_to_guest_tosql")
|
return nil, errors.Wrap(err, "demote_user_to_guest_tosql")
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := transaction.Exec(queryString, args...); err != nil {
|
if _, err := transaction.Exec(queryString, args...); err != nil {
|
||||||
return errors.Wrapf(err, "failed to update TeamMembers with userId=%s", userId)
|
return nil, errors.Wrapf(err, "failed to update TeamMembers with userId=%s", userID)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := transaction.Commit(); err != nil {
|
if err := transaction.Commit(); err != nil {
|
||||||
return errors.Wrap(err, "commit_transaction")
|
return nil, errors.Wrap(err, "commit_transaction")
|
||||||
}
|
}
|
||||||
return nil
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (us SqlUserStore) AutocompleteUsersInChannel(teamId, channelId, term string, options *model.UserSearchOptions) (*model.UserAutocompleteInChannel, error) {
|
func (us SqlUserStore) AutocompleteUsersInChannel(teamId, channelId, term string, options *model.UserSearchOptions) (*model.UserAutocompleteInChannel, error) {
|
||||||
|
|||||||
@@ -322,21 +322,21 @@ type UserStore interface {
|
|||||||
UpdateAuthData(userId string, service string, authData *string, email string, resetMfa bool) (string, error)
|
UpdateAuthData(userId string, service string, authData *string, email string, resetMfa bool) (string, error)
|
||||||
UpdateMfaSecret(userId, secret string) error
|
UpdateMfaSecret(userId, secret string) error
|
||||||
UpdateMfaActive(userId string, active bool) error
|
UpdateMfaActive(userId string, active bool) error
|
||||||
Get(id string) (*model.User, error)
|
Get(ctx context.Context, id string) (*model.User, error)
|
||||||
GetMany(ids []string) ([]*model.User, error)
|
GetMany(ctx context.Context, ids []string) ([]*model.User, error)
|
||||||
GetAll() ([]*model.User, error)
|
GetAll() ([]*model.User, error)
|
||||||
ClearCaches()
|
ClearCaches()
|
||||||
InvalidateProfilesInChannelCacheByUser(userId string)
|
InvalidateProfilesInChannelCacheByUser(userId string)
|
||||||
InvalidateProfilesInChannelCache(channelId string)
|
InvalidateProfilesInChannelCache(channelId string)
|
||||||
GetProfilesInChannel(options *model.UserGetOptions) ([]*model.User, error)
|
GetProfilesInChannel(options *model.UserGetOptions) ([]*model.User, error)
|
||||||
GetProfilesInChannelByStatus(options *model.UserGetOptions) ([]*model.User, error)
|
GetProfilesInChannelByStatus(options *model.UserGetOptions) ([]*model.User, error)
|
||||||
GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error)
|
GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error)
|
||||||
GetProfilesNotInChannel(teamId string, channelId string, groupConstrained bool, offset int, limit int, viewRestrictions *model.ViewUsersRestrictions) ([]*model.User, error)
|
GetProfilesNotInChannel(teamId string, channelId string, groupConstrained bool, offset int, limit int, viewRestrictions *model.ViewUsersRestrictions) ([]*model.User, error)
|
||||||
GetProfilesWithoutTeam(options *model.UserGetOptions) ([]*model.User, error)
|
GetProfilesWithoutTeam(options *model.UserGetOptions) ([]*model.User, error)
|
||||||
GetProfilesByUsernames(usernames []string, viewRestrictions *model.ViewUsersRestrictions) ([]*model.User, error)
|
GetProfilesByUsernames(usernames []string, viewRestrictions *model.ViewUsersRestrictions) ([]*model.User, error)
|
||||||
GetAllProfiles(options *model.UserGetOptions) ([]*model.User, error)
|
GetAllProfiles(options *model.UserGetOptions) ([]*model.User, error)
|
||||||
GetProfiles(options *model.UserGetOptions) ([]*model.User, error)
|
GetProfiles(options *model.UserGetOptions) ([]*model.User, error)
|
||||||
GetProfileByIds(userIds []string, options *UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error)
|
GetProfileByIds(ctx context.Context, userIds []string, options *UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error)
|
||||||
GetProfileByGroupChannelIdsForUser(userId string, channelIds []string) (map[string][]*model.User, error)
|
GetProfileByGroupChannelIdsForUser(userId string, channelIds []string) (map[string][]*model.User, error)
|
||||||
InvalidateProfileCacheForUser(userId string)
|
InvalidateProfileCacheForUser(userId string)
|
||||||
GetByEmail(email string) (*model.User, error)
|
GetByEmail(email string) (*model.User, error)
|
||||||
@@ -378,7 +378,7 @@ type UserStore interface {
|
|||||||
GetTeamGroupUsers(teamID string) ([]*model.User, error)
|
GetTeamGroupUsers(teamID string) ([]*model.User, error)
|
||||||
GetChannelGroupUsers(channelID string) ([]*model.User, error)
|
GetChannelGroupUsers(channelID string) ([]*model.User, error)
|
||||||
PromoteGuestToUser(userID string) error
|
PromoteGuestToUser(userID string) error
|
||||||
DemoteUserToGuest(userID string) error
|
DemoteUserToGuest(userID string) (*model.User, error)
|
||||||
DeactivateGuests() ([]string, error)
|
DeactivateGuests() ([]string, error)
|
||||||
AutocompleteUsersInChannel(teamId, channelId, term string, options *model.UserSearchOptions) (*model.UserAutocompleteInChannel, error)
|
AutocompleteUsersInChannel(teamId, channelId, term string, options *model.UserSearchOptions) (*model.UserAutocompleteInChannel, error)
|
||||||
GetKnownUsers(userID string) ([]string, error)
|
GetKnownUsers(userID string) ([]string, error)
|
||||||
|
|||||||
@@ -5,9 +5,12 @@
|
|||||||
package mocks
|
package mocks
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
context "context"
|
||||||
|
|
||||||
model "github.com/mattermost/mattermost-server/v5/model"
|
model "github.com/mattermost/mattermost-server/v5/model"
|
||||||
store "github.com/mattermost/mattermost-server/v5/store"
|
|
||||||
mock "github.com/stretchr/testify/mock"
|
mock "github.com/stretchr/testify/mock"
|
||||||
|
|
||||||
|
store "github.com/mattermost/mattermost-server/v5/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
// UserStore is an autogenerated mock type for the UserStore type
|
// UserStore is an autogenerated mock type for the UserStore type
|
||||||
@@ -228,26 +231,12 @@ func (_m *UserStore) DeactivateGuests() ([]string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// DemoteUserToGuest provides a mock function with given fields: userID
|
// DemoteUserToGuest provides a mock function with given fields: userID
|
||||||
func (_m *UserStore) DemoteUserToGuest(userID string) error {
|
func (_m *UserStore) DemoteUserToGuest(userID string) (*model.User, error) {
|
||||||
ret := _m.Called(userID)
|
ret := _m.Called(userID)
|
||||||
|
|
||||||
var r0 error
|
|
||||||
if rf, ok := ret.Get(0).(func(string) error); ok {
|
|
||||||
r0 = rf(userID)
|
|
||||||
} else {
|
|
||||||
r0 = ret.Error(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
return r0
|
|
||||||
}
|
|
||||||
|
|
||||||
// Get provides a mock function with given fields: id
|
|
||||||
func (_m *UserStore) Get(id string) (*model.User, error) {
|
|
||||||
ret := _m.Called(id)
|
|
||||||
|
|
||||||
var r0 *model.User
|
var r0 *model.User
|
||||||
if rf, ok := ret.Get(0).(func(string) *model.User); ok {
|
if rf, ok := ret.Get(0).(func(string) *model.User); ok {
|
||||||
r0 = rf(id)
|
r0 = rf(userID)
|
||||||
} else {
|
} else {
|
||||||
if ret.Get(0) != nil {
|
if ret.Get(0) != nil {
|
||||||
r0 = ret.Get(0).(*model.User)
|
r0 = ret.Get(0).(*model.User)
|
||||||
@@ -256,7 +245,30 @@ func (_m *UserStore) Get(id string) (*model.User, error) {
|
|||||||
|
|
||||||
var r1 error
|
var r1 error
|
||||||
if rf, ok := ret.Get(1).(func(string) error); ok {
|
if rf, ok := ret.Get(1).(func(string) error); ok {
|
||||||
r1 = rf(id)
|
r1 = rf(userID)
|
||||||
|
} else {
|
||||||
|
r1 = ret.Error(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
return r0, r1
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get provides a mock function with given fields: ctx, id
|
||||||
|
func (_m *UserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
|
ret := _m.Called(ctx, id)
|
||||||
|
|
||||||
|
var r0 *model.User
|
||||||
|
if rf, ok := ret.Get(0).(func(context.Context, string) *model.User); ok {
|
||||||
|
r0 = rf(ctx, id)
|
||||||
|
} else {
|
||||||
|
if ret.Get(0) != nil {
|
||||||
|
r0 = ret.Get(0).(*model.User)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var r1 error
|
||||||
|
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
|
||||||
|
r1 = rf(ctx, id)
|
||||||
} else {
|
} else {
|
||||||
r1 = ret.Error(1)
|
r1 = ret.Error(1)
|
||||||
}
|
}
|
||||||
@@ -356,13 +368,13 @@ func (_m *UserStore) GetAllProfiles(options *model.UserGetOptions) ([]*model.Use
|
|||||||
return r0, r1
|
return r0, r1
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetAllProfilesInChannel provides a mock function with given fields: channelId, allowFromCache
|
// GetAllProfilesInChannel provides a mock function with given fields: ctx, channelId, allowFromCache
|
||||||
func (_m *UserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (_m *UserStore) GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
ret := _m.Called(channelId, allowFromCache)
|
ret := _m.Called(ctx, channelId, allowFromCache)
|
||||||
|
|
||||||
var r0 map[string]*model.User
|
var r0 map[string]*model.User
|
||||||
if rf, ok := ret.Get(0).(func(string, bool) map[string]*model.User); ok {
|
if rf, ok := ret.Get(0).(func(context.Context, string, bool) map[string]*model.User); ok {
|
||||||
r0 = rf(channelId, allowFromCache)
|
r0 = rf(ctx, channelId, allowFromCache)
|
||||||
} else {
|
} else {
|
||||||
if ret.Get(0) != nil {
|
if ret.Get(0) != nil {
|
||||||
r0 = ret.Get(0).(map[string]*model.User)
|
r0 = ret.Get(0).(map[string]*model.User)
|
||||||
@@ -370,8 +382,8 @@ func (_m *UserStore) GetAllProfilesInChannel(channelId string, allowFromCache bo
|
|||||||
}
|
}
|
||||||
|
|
||||||
var r1 error
|
var r1 error
|
||||||
if rf, ok := ret.Get(1).(func(string, bool) error); ok {
|
if rf, ok := ret.Get(1).(func(context.Context, string, bool) error); ok {
|
||||||
r1 = rf(channelId, allowFromCache)
|
r1 = rf(ctx, channelId, allowFromCache)
|
||||||
} else {
|
} else {
|
||||||
r1 = ret.Error(1)
|
r1 = ret.Error(1)
|
||||||
}
|
}
|
||||||
@@ -603,13 +615,13 @@ func (_m *UserStore) GetKnownUsers(userID string) ([]string, error) {
|
|||||||
return r0, r1
|
return r0, r1
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetMany provides a mock function with given fields: ids
|
// GetMany provides a mock function with given fields: ctx, ids
|
||||||
func (_m *UserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (_m *UserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
ret := _m.Called(ids)
|
ret := _m.Called(ctx, ids)
|
||||||
|
|
||||||
var r0 []*model.User
|
var r0 []*model.User
|
||||||
if rf, ok := ret.Get(0).(func([]string) []*model.User); ok {
|
if rf, ok := ret.Get(0).(func(context.Context, []string) []*model.User); ok {
|
||||||
r0 = rf(ids)
|
r0 = rf(ctx, ids)
|
||||||
} else {
|
} else {
|
||||||
if ret.Get(0) != nil {
|
if ret.Get(0) != nil {
|
||||||
r0 = ret.Get(0).([]*model.User)
|
r0 = ret.Get(0).([]*model.User)
|
||||||
@@ -617,8 +629,8 @@ func (_m *UserStore) GetMany(ids []string) ([]*model.User, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var r1 error
|
var r1 error
|
||||||
if rf, ok := ret.Get(1).(func([]string) error); ok {
|
if rf, ok := ret.Get(1).(func(context.Context, []string) error); ok {
|
||||||
r1 = rf(ids)
|
r1 = rf(ctx, ids)
|
||||||
} else {
|
} else {
|
||||||
r1 = ret.Error(1)
|
r1 = ret.Error(1)
|
||||||
}
|
}
|
||||||
@@ -672,13 +684,13 @@ func (_m *UserStore) GetProfileByGroupChannelIdsForUser(userId string, channelId
|
|||||||
return r0, r1
|
return r0, r1
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetProfileByIds provides a mock function with given fields: userIds, options, allowFromCache
|
// GetProfileByIds provides a mock function with given fields: ctx, userIds, options, allowFromCache
|
||||||
func (_m *UserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (_m *UserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
ret := _m.Called(userIds, options, allowFromCache)
|
ret := _m.Called(ctx, userIds, options, allowFromCache)
|
||||||
|
|
||||||
var r0 []*model.User
|
var r0 []*model.User
|
||||||
if rf, ok := ret.Get(0).(func([]string, *store.UserGetByIdsOpts, bool) []*model.User); ok {
|
if rf, ok := ret.Get(0).(func(context.Context, []string, *store.UserGetByIdsOpts, bool) []*model.User); ok {
|
||||||
r0 = rf(userIds, options, allowFromCache)
|
r0 = rf(ctx, userIds, options, allowFromCache)
|
||||||
} else {
|
} else {
|
||||||
if ret.Get(0) != nil {
|
if ret.Get(0) != nil {
|
||||||
r0 = ret.Get(0).([]*model.User)
|
r0 = ret.Get(0).([]*model.User)
|
||||||
@@ -686,8 +698,8 @@ func (_m *UserStore) GetProfileByIds(userIds []string, options *store.UserGetByI
|
|||||||
}
|
}
|
||||||
|
|
||||||
var r1 error
|
var r1 error
|
||||||
if rf, ok := ret.Get(1).(func([]string, *store.UserGetByIdsOpts, bool) error); ok {
|
if rf, ok := ret.Get(1).(func(context.Context, []string, *store.UserGetByIdsOpts, bool) error); ok {
|
||||||
r1 = rf(userIds, options, allowFromCache)
|
r1 = rf(ctx, userIds, options, allowFromCache)
|
||||||
} else {
|
} else {
|
||||||
r1 = ret.Error(1)
|
r1 = ret.Error(1)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2829,7 +2829,7 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) {
|
|||||||
require.Equal(t, maxUsersPerTeam, int(totalMemberCount), "should have 5 team members again, had %v instead", totalMemberCount)
|
require.Equal(t, maxUsersPerTeam, int(totalMemberCount), "should have 5 team members again, had %v instead", totalMemberCount)
|
||||||
|
|
||||||
// Deactivating a user should make them stop counting against max members
|
// Deactivating a user should make them stop counting against max members
|
||||||
user2, nErr := ss.User().Get(userIds[1])
|
user2, nErr := ss.User().Get(context.Background(), userIds[1])
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
user2.DeleteAt = 1234
|
user2.DeleteAt = 1234
|
||||||
_, nErr = ss.User().Update(user2, true)
|
_, nErr = ss.User().Update(user2, true)
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package storetest
|
package storetest
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -226,7 +227,7 @@ func testUserStoreUpdateUpdateAt(t *testing.T, ss store.Store) {
|
|||||||
_, err = ss.User().UpdateUpdateAt(u1.Id)
|
_, err = ss.User().UpdateUpdateAt(u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
user, err := ss.User().Get(u1.Id)
|
user, err := ss.User().Get(context.Background(), u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Less(t, u1.UpdateAt, user.UpdateAt, "UpdateAt not updated correctly")
|
require.Less(t, u1.UpdateAt, user.UpdateAt, "UpdateAt not updated correctly")
|
||||||
}
|
}
|
||||||
@@ -243,7 +244,7 @@ func testUserStoreUpdateFailedPasswordAttempts(t *testing.T, ss store.Store) {
|
|||||||
err = ss.User().UpdateFailedPasswordAttempts(u1.Id, 3)
|
err = ss.User().UpdateFailedPasswordAttempts(u1.Id, 3)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
user, err := ss.User().Get(u1.Id)
|
user, err := ss.User().Get(context.Background(), u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, 3, user.FailedAttempts, "FailedAttempts not updated correctly")
|
require.Equal(t, 3, user.FailedAttempts, "FailedAttempts not updated correctly")
|
||||||
}
|
}
|
||||||
@@ -276,19 +277,19 @@ func testUserStoreGet(t *testing.T, ss store.Store) {
|
|||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
t.Run("fetch empty id", func(t *testing.T) {
|
t.Run("fetch empty id", func(t *testing.T) {
|
||||||
_, err := ss.User().Get("")
|
_, err := ss.User().Get(context.Background(), "")
|
||||||
require.Error(t, err)
|
require.Error(t, err)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("fetch user 1", func(t *testing.T) {
|
t.Run("fetch user 1", func(t *testing.T) {
|
||||||
actual, err := ss.User().Get(u1.Id)
|
actual, err := ss.User().Get(context.Background(), u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, u1, actual)
|
require.Equal(t, u1, actual)
|
||||||
require.False(t, actual.IsBot)
|
require.False(t, actual.IsBot)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("fetch user 2, also a bot", func(t *testing.T) {
|
t.Run("fetch user 2, also a bot", func(t *testing.T) {
|
||||||
actual, err := ss.User().Get(u2.Id)
|
actual, err := ss.User().Get(context.Background(), u2.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, u2, actual)
|
require.Equal(t, u2, actual)
|
||||||
require.True(t, actual.IsBot)
|
require.True(t, actual.IsBot)
|
||||||
@@ -1272,7 +1273,7 @@ func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
t.Run("all profiles in channel 1, no caching", func(t *testing.T) {
|
t.Run("all profiles in channel 1, no caching", func(t *testing.T) {
|
||||||
var profiles map[string]*model.User
|
var profiles map[string]*model.User
|
||||||
profiles, err = ss.User().GetAllProfilesInChannel(c1.Id, false)
|
profiles, err = ss.User().GetAllProfilesInChannel(context.Background(), c1.Id, false)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, map[string]*model.User{
|
assert.Equal(t, map[string]*model.User{
|
||||||
u1.Id: sanitized(u1),
|
u1.Id: sanitized(u1),
|
||||||
@@ -1283,7 +1284,7 @@ func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
t.Run("all profiles in channel 2, no caching", func(t *testing.T) {
|
t.Run("all profiles in channel 2, no caching", func(t *testing.T) {
|
||||||
var profiles map[string]*model.User
|
var profiles map[string]*model.User
|
||||||
profiles, err = ss.User().GetAllProfilesInChannel(c2.Id, false)
|
profiles, err = ss.User().GetAllProfilesInChannel(context.Background(), c2.Id, false)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, map[string]*model.User{
|
assert.Equal(t, map[string]*model.User{
|
||||||
u1.Id: sanitized(u1),
|
u1.Id: sanitized(u1),
|
||||||
@@ -1292,7 +1293,7 @@ func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
t.Run("all profiles in channel 2, caching", func(t *testing.T) {
|
t.Run("all profiles in channel 2, caching", func(t *testing.T) {
|
||||||
var profiles map[string]*model.User
|
var profiles map[string]*model.User
|
||||||
profiles, err = ss.User().GetAllProfilesInChannel(c2.Id, true)
|
profiles, err = ss.User().GetAllProfilesInChannel(context.Background(), c2.Id, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, map[string]*model.User{
|
assert.Equal(t, map[string]*model.User{
|
||||||
u1.Id: sanitized(u1),
|
u1.Id: sanitized(u1),
|
||||||
@@ -1301,7 +1302,7 @@ func testUserStoreGetAllProfilesInChannel(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
t.Run("all profiles in channel 2, caching [repeated]", func(t *testing.T) {
|
t.Run("all profiles in channel 2, caching [repeated]", func(t *testing.T) {
|
||||||
var profiles map[string]*model.User
|
var profiles map[string]*model.User
|
||||||
profiles, err = ss.User().GetAllProfilesInChannel(c2.Id, true)
|
profiles, err = ss.User().GetAllProfilesInChannel(context.Background(), c2.Id, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, map[string]*model.User{
|
assert.Equal(t, map[string]*model.User{
|
||||||
u1.Id: sanitized(u1),
|
u1.Id: sanitized(u1),
|
||||||
@@ -1521,37 +1522,37 @@ func testUserStoreGetProfilesByIds(t *testing.T, ss store.Store) {
|
|||||||
defer func() { require.NoError(t, ss.User().PermanentDelete(u4.Id)) }()
|
defer func() { require.NoError(t, ss.User().PermanentDelete(u4.Id)) }()
|
||||||
|
|
||||||
t.Run("get u1 by id, no caching", func(t *testing.T) {
|
t.Run("get u1 by id, no caching", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{u1.Id}, nil, false)
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{u1.Id}, nil, false)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []*model.User{u1}, users)
|
assert.Equal(t, []*model.User{u1}, users)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("get u1 by id, caching", func(t *testing.T) {
|
t.Run("get u1 by id, caching", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{u1.Id}, nil, true)
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{u1.Id}, nil, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []*model.User{u1}, users)
|
assert.Equal(t, []*model.User{u1}, users)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("get u1, u2, u3 by id, no caching", func(t *testing.T) {
|
t.Run("get u1, u2, u3 by id, no caching", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{u1.Id, u2.Id, u3.Id}, nil, false)
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{u1.Id, u2.Id, u3.Id}, nil, false)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []*model.User{u1, u2, u3}, users)
|
assert.Equal(t, []*model.User{u1, u2, u3}, users)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("get u1, u2, u3 by id, caching", func(t *testing.T) {
|
t.Run("get u1, u2, u3 by id, caching", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{u1.Id, u2.Id, u3.Id}, nil, true)
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{u1.Id, u2.Id, u3.Id}, nil, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []*model.User{u1, u2, u3}, users)
|
assert.Equal(t, []*model.User{u1, u2, u3}, users)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("get unknown id, caching", func(t *testing.T) {
|
t.Run("get unknown id, caching", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{"123"}, nil, true)
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{"123"}, nil, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, []*model.User{}, users)
|
assert.Equal(t, []*model.User{}, users)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("should only return users with UpdateAt greater than the since time", func(t *testing.T) {
|
t.Run("should only return users with UpdateAt greater than the since time", func(t *testing.T) {
|
||||||
users, err := ss.User().GetProfileByIds([]string{u1.Id, u2.Id, u3.Id, u4.Id}, &store.UserGetByIdsOpts{
|
users, err := ss.User().GetProfileByIds(context.Background(), []string{u1.Id, u2.Id, u3.Id, u4.Id}, &store.UserGetByIdsOpts{
|
||||||
Since: u2.CreateAt,
|
Since: u2.CreateAt,
|
||||||
}, true)
|
}, true)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -4835,7 +4836,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", updatedUser.Roles)
|
require.Equal(t, "system_user", updatedUser.Roles)
|
||||||
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
|
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
|
||||||
@@ -4881,7 +4882,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user system_admin", updatedUser.Roles)
|
require.Equal(t, "system_user system_admin", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -4912,7 +4913,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", updatedUser.Roles)
|
require.Equal(t, "system_user", updatedUser.Roles)
|
||||||
})
|
})
|
||||||
@@ -4937,7 +4938,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", updatedUser.Roles)
|
require.Equal(t, "system_user", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -4977,7 +4978,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", updatedUser.Roles)
|
require.Equal(t, "system_user", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5022,7 +5023,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user.Id)
|
err = ss.User().PromoteGuestToUser(user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user custom_role", updatedUser.Roles)
|
require.Equal(t, "system_user custom_role", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5088,7 +5089,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
err = ss.User().PromoteGuestToUser(user1.Id)
|
err = ss.User().PromoteGuestToUser(user1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
updatedUser, err := ss.User().Get(user1.Id)
|
updatedUser, err := ss.User().Get(context.Background(), user1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", updatedUser.Roles)
|
require.Equal(t, "system_user", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5102,7 +5103,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
|
|||||||
require.False(t, updatedChannelMember.SchemeGuest)
|
require.False(t, updatedChannelMember.SchemeGuest)
|
||||||
require.True(t, updatedChannelMember.SchemeUser)
|
require.True(t, updatedChannelMember.SchemeUser)
|
||||||
|
|
||||||
notUpdatedUser, err := ss.User().Get(user2.Id)
|
notUpdatedUser, err := ss.User().Get(context.Background(), user2.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", notUpdatedUser.Roles)
|
require.Equal(t, "system_guest", notUpdatedUser.Roles)
|
||||||
|
|
||||||
@@ -5148,19 +5149,17 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
|
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
|
||||||
|
|
||||||
updatedTeamMember, nErr := ss.Team().GetMember(teamId, user.Id)
|
updatedTeamMember, nErr := ss.Team().GetMember(teamId, updatedUser.Id)
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
require.True(t, updatedTeamMember.SchemeGuest)
|
require.True(t, updatedTeamMember.SchemeGuest)
|
||||||
require.False(t, updatedTeamMember.SchemeUser)
|
require.False(t, updatedTeamMember.SchemeUser)
|
||||||
|
|
||||||
updatedChannelMember, nErr := ss.Channel().GetMember(channel.Id, user.Id)
|
updatedChannelMember, nErr := ss.Channel().GetMember(channel.Id, updatedUser.Id)
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
require.True(t, updatedChannelMember.SchemeGuest)
|
require.True(t, updatedChannelMember.SchemeGuest)
|
||||||
require.False(t, updatedChannelMember.SchemeUser)
|
require.False(t, updatedChannelMember.SchemeUser)
|
||||||
@@ -5194,9 +5193,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: true, SchemeUser: false, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: true, SchemeUser: false, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5225,9 +5222,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer func() { require.NoError(t, ss.User().PermanentDelete(user.Id)) }()
|
defer func() { require.NoError(t, ss.User().PermanentDelete(user.Id)) }()
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
})
|
})
|
||||||
@@ -5250,9 +5245,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr := ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: user.Id, SchemeGuest: false, SchemeUser: true}, 999)
|
_, nErr := ss.Team().SaveMember(&model.TeamMember{TeamId: teamId, UserId: user.Id, SchemeGuest: false, SchemeUser: true}, 999)
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5290,9 +5283,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5335,9 +5326,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest custom_role", updatedUser.Roles)
|
require.Equal(t, "system_guest custom_role", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5401,9 +5390,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user2.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
_, nErr = ss.Channel().SaveMember(&model.ChannelMember{ChannelId: channel.Id, UserId: user2.Id, SchemeGuest: false, SchemeUser: true, NotifyProps: model.GetDefaultChannelNotifyProps()})
|
||||||
require.NoError(t, nErr)
|
require.NoError(t, nErr)
|
||||||
|
|
||||||
err = ss.User().DemoteUserToGuest(user1.Id)
|
updatedUser, err := ss.User().DemoteUserToGuest(user1.Id)
|
||||||
require.NoError(t, err)
|
|
||||||
updatedUser, err := ss.User().Get(user1.Id)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_guest", updatedUser.Roles)
|
require.Equal(t, "system_guest", updatedUser.Roles)
|
||||||
|
|
||||||
@@ -5417,7 +5404,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
|
|||||||
require.True(t, updatedChannelMember.SchemeGuest)
|
require.True(t, updatedChannelMember.SchemeGuest)
|
||||||
require.False(t, updatedChannelMember.SchemeUser)
|
require.False(t, updatedChannelMember.SchemeUser)
|
||||||
|
|
||||||
notUpdatedUser, err := ss.User().Get(user2.Id)
|
notUpdatedUser, err := ss.User().Get(context.Background(), user2.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, "system_user", notUpdatedUser.Roles)
|
require.Equal(t, "system_user", notUpdatedUser.Roles)
|
||||||
|
|
||||||
@@ -5493,19 +5480,19 @@ func testDeactivateGuests(t *testing.T, ss store.Store) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.ElementsMatch(t, []string{guest1.Id, guest2.Id}, ids)
|
assert.ElementsMatch(t, []string{guest1.Id, guest2.Id}, ids)
|
||||||
|
|
||||||
u, err := ss.User().Get(guest1.Id)
|
u, err := ss.User().Get(context.Background(), guest1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.NotEqual(t, u.DeleteAt, int64(0))
|
assert.NotEqual(t, u.DeleteAt, int64(0))
|
||||||
|
|
||||||
u, err = ss.User().Get(guest2.Id)
|
u, err = ss.User().Get(context.Background(), guest2.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.NotEqual(t, u.DeleteAt, int64(0))
|
assert.NotEqual(t, u.DeleteAt, int64(0))
|
||||||
|
|
||||||
u, err = ss.User().Get(guest3.Id)
|
u, err = ss.User().Get(context.Background(), guest3.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, u.DeleteAt, int64(10))
|
assert.Equal(t, u.DeleteAt, int64(10))
|
||||||
|
|
||||||
u, err = ss.User().Get(regularUser.Id)
|
u, err = ss.User().Get(context.Background(), regularUser.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, u.DeleteAt, int64(0))
|
assert.Equal(t, u.DeleteAt, int64(0))
|
||||||
})
|
})
|
||||||
@@ -5523,7 +5510,7 @@ func testUserStoreResetLastPictureUpdate(t *testing.T, ss store.Store) {
|
|||||||
err = ss.User().UpdateLastPictureUpdate(u1.Id)
|
err = ss.User().UpdateLastPictureUpdate(u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
user, err := ss.User().Get(u1.Id)
|
user, err := ss.User().Get(context.Background(), u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
assert.NotZero(t, user.LastPictureUpdate)
|
assert.NotZero(t, user.LastPictureUpdate)
|
||||||
@@ -5537,7 +5524,7 @@ func testUserStoreResetLastPictureUpdate(t *testing.T, ss store.Store) {
|
|||||||
|
|
||||||
ss.User().InvalidateProfileCacheForUser(u1.Id)
|
ss.User().InvalidateProfileCacheForUser(u1.Id)
|
||||||
|
|
||||||
user2, err := ss.User().Get(u1.Id)
|
user2, err := ss.User().Get(context.Background(), u1.Id)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
assert.True(t, user2.UpdateAt > user.UpdateAt)
|
assert.True(t, user2.UpdateAt > user.UpdateAt)
|
||||||
|
|||||||
@@ -7574,10 +7574,10 @@ func (s *TimerLayerUserStore) DeactivateGuests() ([]string, error) {
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TimerLayerUserStore) DemoteUserToGuest(userID string) error {
|
func (s *TimerLayerUserStore) DemoteUserToGuest(userID string) (*model.User, error) {
|
||||||
start := timemodule.Now()
|
start := timemodule.Now()
|
||||||
|
|
||||||
err := s.UserStore.DemoteUserToGuest(userID)
|
result, err := s.UserStore.DemoteUserToGuest(userID)
|
||||||
|
|
||||||
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
||||||
if s.Root.Metrics != nil {
|
if s.Root.Metrics != nil {
|
||||||
@@ -7587,13 +7587,13 @@ func (s *TimerLayerUserStore) DemoteUserToGuest(userID string) error {
|
|||||||
}
|
}
|
||||||
s.Root.Metrics.ObserveStoreMethodDuration("UserStore.DemoteUserToGuest", success, elapsed)
|
s.Root.Metrics.ObserveStoreMethodDuration("UserStore.DemoteUserToGuest", success, elapsed)
|
||||||
}
|
}
|
||||||
return err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TimerLayerUserStore) Get(id string) (*model.User, error) {
|
func (s *TimerLayerUserStore) Get(ctx context.Context, id string) (*model.User, error) {
|
||||||
start := timemodule.Now()
|
start := timemodule.Now()
|
||||||
|
|
||||||
result, err := s.UserStore.Get(id)
|
result, err := s.UserStore.Get(ctx, id)
|
||||||
|
|
||||||
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
||||||
if s.Root.Metrics != nil {
|
if s.Root.Metrics != nil {
|
||||||
@@ -7670,10 +7670,10 @@ func (s *TimerLayerUserStore) GetAllProfiles(options *model.UserGetOptions) ([]*
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TimerLayerUserStore) GetAllProfilesInChannel(channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
func (s *TimerLayerUserStore) GetAllProfilesInChannel(ctx context.Context, channelId string, allowFromCache bool) (map[string]*model.User, error) {
|
||||||
start := timemodule.Now()
|
start := timemodule.Now()
|
||||||
|
|
||||||
result, err := s.UserStore.GetAllProfilesInChannel(channelId, allowFromCache)
|
result, err := s.UserStore.GetAllProfilesInChannel(ctx, channelId, allowFromCache)
|
||||||
|
|
||||||
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
||||||
if s.Root.Metrics != nil {
|
if s.Root.Metrics != nil {
|
||||||
@@ -7862,10 +7862,10 @@ func (s *TimerLayerUserStore) GetKnownUsers(userID string) ([]string, error) {
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TimerLayerUserStore) GetMany(ids []string) ([]*model.User, error) {
|
func (s *TimerLayerUserStore) GetMany(ctx context.Context, ids []string) ([]*model.User, error) {
|
||||||
start := timemodule.Now()
|
start := timemodule.Now()
|
||||||
|
|
||||||
result, err := s.UserStore.GetMany(ids)
|
result, err := s.UserStore.GetMany(ctx, ids)
|
||||||
|
|
||||||
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
||||||
if s.Root.Metrics != nil {
|
if s.Root.Metrics != nil {
|
||||||
@@ -7910,10 +7910,10 @@ func (s *TimerLayerUserStore) GetProfileByGroupChannelIdsForUser(userId string,
|
|||||||
return result, err
|
return result, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TimerLayerUserStore) GetProfileByIds(userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
func (s *TimerLayerUserStore) GetProfileByIds(ctx context.Context, userIds []string, options *store.UserGetByIdsOpts, allowFromCache bool) ([]*model.User, error) {
|
||||||
start := timemodule.Now()
|
start := timemodule.Now()
|
||||||
|
|
||||||
result, err := s.UserStore.GetProfileByIds(userIds, options, allowFromCache)
|
result, err := s.UserStore.GetProfileByIds(ctx, userIds, options, allowFromCache)
|
||||||
|
|
||||||
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
|
||||||
if s.Root.Metrics != nil {
|
if s.Root.Metrics != nil {
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
package web
|
package web
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"net/http"
|
"net/http"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -55,7 +56,7 @@ func TestMfaRequired(t *testing.T) {
|
|||||||
mockStore := th.App.Srv().Store.(*mocks.Store)
|
mockStore := th.App.Srv().Store.(*mocks.Store)
|
||||||
mockUserStore := mocks.UserStore{}
|
mockUserStore := mocks.UserStore{}
|
||||||
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
mockUserStore.On("Count", mock.Anything).Return(int64(10), nil)
|
||||||
mockUserStore.On("Get", "userid").Return(nil, model.NewAppError("Userstore.Get", "storeerror", nil, "store error", http.StatusInternalServerError))
|
mockUserStore.On("Get", context.Background(), "userid").Return(nil, model.NewAppError("Userstore.Get", "storeerror", nil, "store error", http.StatusInternalServerError))
|
||||||
mockPostStore := mocks.PostStore{}
|
mockPostStore := mocks.PostStore{}
|
||||||
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
mockPostStore.On("GetMaxPostSize").Return(65535, nil)
|
||||||
mockSystemStore := mocks.SystemStore{}
|
mockSystemStore := mocks.SystemStore{}
|
||||||
|
|||||||
Ссылка в новой задаче
Block a user