From 6fd174a95fa68fc183b112096d0dec70a04288cb Mon Sep 17 00:00:00 2001 From: Claudio Costa Date: Tue, 13 Dec 2022 11:47:05 -0600 Subject: [PATCH] [MM-48946] Fix read after write issue when uploading data (#21868) * Fix read after write issue when uploading data * Prefer request.CTX interface --- api4/remote_cluster.go | 4 +++- api4/upload.go | 6 ++++-- app/app_iface.go | 4 ++-- app/opentracing/opentracing_layer.go | 6 +++--- app/plugin_api.go | 4 +++- app/upload.go | 11 ++++++----- store/opentracinglayer/opentracinglayer.go | 4 ++-- store/retrylayer/retrylayer.go | 4 ++-- store/sqlstore/upload_session_store.go | 5 +++-- store/store.go | 2 +- store/storetest/mocks/UploadSessionStore.go | 16 +++++++++------- store/storetest/upload_session_store.go | 9 +++++---- store/timerlayer/timerlayer.go | 4 ++-- 13 files changed, 45 insertions(+), 34 deletions(-) diff --git a/api4/remote_cluster.go b/api4/remote_cluster.go index 0fe0be57be..c70d767cd3 100644 --- a/api4/remote_cluster.go +++ b/api4/remote_cluster.go @@ -9,6 +9,7 @@ import ( "net/http" "time" + "github.com/mattermost/mattermost-server/v6/app" "github.com/mattermost/mattermost-server/v6/audit" "github.com/mattermost/mattermost-server/v6/model" "github.com/mattermost/mattermost-server/v6/services/remotecluster" @@ -194,7 +195,8 @@ func uploadRemoteData(c *Context, w http.ResponseWriter, r *http.Request) { defer c.LogAuditRec(auditRec) auditRec.AddEventParameter("upload_id", c.Params.UploadId) - us, err := c.App.GetUploadSession(c.Params.UploadId) + c.AppContext.SetContext(app.WithMaster(c.AppContext.Context())) + us, err := c.App.GetUploadSession(c.AppContext, c.Params.UploadId) if err != nil { c.Err = err return diff --git a/api4/upload.go b/api4/upload.go index a89bc1149c..37f5a1bb01 100644 --- a/api4/upload.go +++ b/api4/upload.go @@ -10,6 +10,7 @@ import ( "mime/multipart" "net/http" + "github.com/mattermost/mattermost-server/v6/app" "github.com/mattermost/mattermost-server/v6/audit" "github.com/mattermost/mattermost-server/v6/model" "github.com/mattermost/mattermost-server/v6/shared/mlog" @@ -91,7 +92,7 @@ func getUpload(c *Context, w http.ResponseWriter, r *http.Request) { return } - us, err := c.App.GetUploadSession(c.Params.UploadId) + us, err := c.App.GetUploadSession(c.AppContext, c.Params.UploadId) if err != nil { c.Err = err return @@ -123,7 +124,8 @@ func uploadData(c *Context, w http.ResponseWriter, r *http.Request) { defer c.LogAuditRec(auditRec) auditRec.AddEventParameter("upload_id", c.Params.UploadId) - us, err := c.App.GetUploadSession(c.Params.UploadId) + c.AppContext.SetContext(app.WithMaster(c.AppContext.Context())) + us, err := c.App.GetUploadSession(c.AppContext, c.Params.UploadId) if err != nil { c.Err = err return diff --git a/app/app_iface.go b/app/app_iface.go index c27ea210c1..0465d8f916 100644 --- a/app/app_iface.go +++ b/app/app_iface.go @@ -809,7 +809,7 @@ type AppIface interface { GetTopReactionsForUserSince(userID string, teamID string, opts *model.InsightsOpts) (*model.TopReactionList, *model.AppError) GetTopThreadsForTeamSince(c request.CTX, teamID, userID string, opts *model.InsightsOpts) (*model.TopThreadList, *model.AppError) GetTopThreadsForUserSince(c request.CTX, teamID, userID string, opts *model.InsightsOpts) (*model.TopThreadList, *model.AppError) - GetUploadSession(uploadId string) (*model.UploadSession, *model.AppError) + GetUploadSession(c request.CTX, uploadId string) (*model.UploadSession, *model.AppError) GetUploadSessionsForUser(userID string) ([]*model.UploadSession, *model.AppError) GetUser(userID string) (*model.User, *model.AppError) GetUserAccessToken(tokenID string, sanitize bool) (*model.UserAccessToken, *model.AppError) @@ -1147,7 +1147,7 @@ type AppIface interface { UpdateUserAuth(userID string, userAuth *model.UserAuth) (*model.UserAuth, *model.AppError) UpdateUserRoles(c request.CTX, userID string, newRoles string, sendWebSocketEvent bool) (*model.User, *model.AppError) UpdateUserRolesWithUser(c request.CTX, user *model.User, newRoles string, sendWebSocketEvent bool) (*model.User, *model.AppError) - UploadData(c *request.Context, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) + UploadData(c request.CTX, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) UploadEmojiImage(c request.CTX, id string, imageData *multipart.FileHeader) *model.AppError UpsertDraft(c *request.Context, draft *model.Draft, connectionID string) (*model.Draft, *model.AppError) UpsertGroupMember(groupID string, userID string) (*model.GroupMember, *model.AppError) diff --git a/app/opentracing/opentracing_layer.go b/app/opentracing/opentracing_layer.go index 2d60825ebc..4e2f0fb783 100644 --- a/app/opentracing/opentracing_layer.go +++ b/app/opentracing/opentracing_layer.go @@ -10305,7 +10305,7 @@ func (a *OpenTracingAppLayer) GetTotalUsersStats(viewRestrictions *model.ViewUse return resultVar0, resultVar1 } -func (a *OpenTracingAppLayer) GetUploadSession(uploadId string) (*model.UploadSession, *model.AppError) { +func (a *OpenTracingAppLayer) GetUploadSession(c request.CTX, uploadId string) (*model.UploadSession, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.GetUploadSession") @@ -10317,7 +10317,7 @@ func (a *OpenTracingAppLayer) GetUploadSession(uploadId string) (*model.UploadSe }() defer span.Finish() - resultVar0, resultVar1 := a.app.GetUploadSession(uploadId) + resultVar0, resultVar1 := a.app.GetUploadSession(c, uploadId) if resultVar1 != nil { span.LogFields(spanlog.Error(resultVar1)) @@ -18135,7 +18135,7 @@ func (a *OpenTracingAppLayer) UpdateWebConnUserActivity(session model.Session, a a.app.UpdateWebConnUserActivity(session, activityAt) } -func (a *OpenTracingAppLayer) UploadData(c *request.Context, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) { +func (a *OpenTracingAppLayer) UploadData(c request.CTX, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) { origCtx := a.ctx span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.UploadData") diff --git a/app/plugin_api.go b/app/plugin_api.go index accedf37ba..5c8ee9fad1 100644 --- a/app/plugin_api.go +++ b/app/plugin_api.go @@ -1255,7 +1255,9 @@ func (api *PluginAPI) UploadData(us *model.UploadSession, rd io.Reader) (*model. } func (api *PluginAPI) GetUploadSession(uploadID string) (*model.UploadSession, error) { - fi, err := api.app.GetUploadSession(uploadID) + // We want to fetch from master DB to avoid a potential read-after-write on the plugin side. + api.ctx.SetContext(WithMaster(api.ctx.Context())) + fi, err := api.app.GetUploadSession(api.ctx, uploadID) if err != nil { return nil, err } diff --git a/app/upload.go b/app/upload.go index 4da0b853b6..ef4bff2b71 100644 --- a/app/upload.go +++ b/app/upload.go @@ -48,7 +48,7 @@ func (a *App) genFileInfoFromReader(name string, file io.ReadSeeker, size int64) return info, nil } -func (a *App) runPluginsHook(c *request.Context, info *model.FileInfo, file io.Reader) *model.AppError { +func (a *App) runPluginsHook(c request.CTX, info *model.FileInfo, file io.Reader) *model.AppError { filePath := info.Path // using a pipe to avoid loading the whole file content in memory. r, w := io.Pipe() @@ -154,8 +154,8 @@ func (a *App) CreateUploadSession(c request.CTX, us *model.UploadSession) (*mode return us, nil } -func (a *App) GetUploadSession(uploadId string) (*model.UploadSession, *model.AppError) { - us, err := a.Srv().Store().UploadSession().Get(uploadId) +func (a *App) GetUploadSession(c request.CTX, uploadId string) (*model.UploadSession, *model.AppError) { + us, err := a.Srv().Store().UploadSession().Get(c.Context(), uploadId) if err != nil { var nfErr *store.ErrNotFound switch { @@ -179,7 +179,7 @@ func (a *App) GetUploadSessionsForUser(userID string) ([]*model.UploadSession, * return uss, nil } -func (a *App) UploadData(c *request.Context, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) { +func (a *App) UploadData(c request.CTX, us *model.UploadSession, rd io.Reader) (*model.FileInfo, *model.AppError) { // prevent more than one caller to upload data at the same time for a given upload session. // This is to avoid possible inconsistencies. a.ch.uploadLockMapMut.Lock() @@ -202,7 +202,8 @@ func (a *App) UploadData(c *request.Context, us *model.UploadSession, rd io.Read }() // fetch the session from store to check for inconsistencies. - if storedSession, err := a.GetUploadSession(us.Id); err != nil { + c.SetContext(WithMaster(c.Context())) + if storedSession, err := a.GetUploadSession(c, us.Id); err != nil { return nil, err } else if us.FileOffset != storedSession.FileOffset { return nil, model.NewAppError("UploadData", "app.upload.upload_data.concurrent.app_error", diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index aa41c28ab5..241f3edf50 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -10612,7 +10612,7 @@ func (s *OpenTracingLayerUploadSessionStore) Delete(id string) error { return err } -func (s *OpenTracingLayerUploadSessionStore) Get(id string) (*model.UploadSession, error) { +func (s *OpenTracingLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UploadSessionStore.Get") s.Root.Store.SetContext(newCtx) @@ -10621,7 +10621,7 @@ func (s *OpenTracingLayerUploadSessionStore) Get(id string) (*model.UploadSessio }() defer span.Finish() - result, err := s.UploadSessionStore.Get(id) + result, err := s.UploadSessionStore.Get(ctx, id) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index 5aca1535cb..dec858bba9 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -12126,11 +12126,11 @@ func (s *RetryLayerUploadSessionStore) Delete(id string) error { } -func (s *RetryLayerUploadSessionStore) Get(id string) (*model.UploadSession, error) { +func (s *RetryLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { tries := 0 for { - result, err := s.UploadSessionStore.Get(id) + result, err := s.UploadSessionStore.Get(ctx, id) if err == nil { return result, nil } diff --git a/store/sqlstore/upload_session_store.go b/store/sqlstore/upload_session_store.go index e22b2056f3..3efec8f239 100644 --- a/store/sqlstore/upload_session_store.go +++ b/store/sqlstore/upload_session_store.go @@ -4,6 +4,7 @@ package sqlstore import ( + "context" "database/sql" sq "github.com/mattermost/squirrel" @@ -78,7 +79,7 @@ func (us SqlUploadSessionStore) Update(session *model.UploadSession) error { return nil } -func (us SqlUploadSessionStore) Get(id string) (*model.UploadSession, error) { +func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { if !model.IsValidId(id) { return nil, errors.New("SqlUploadSessionStore.Get: id is not valid") } @@ -91,7 +92,7 @@ func (us SqlUploadSessionStore) Get(id string) (*model.UploadSession, error) { return nil, errors.Wrap(err, "SqlUploadSessionStore.Get: failed to build query") } var session model.UploadSession - if err := us.GetReplicaX().Get(&session, query, args...); err != nil { + if err := us.DBXFromContext(ctx).Get(&session, query, args...); err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("UploadSession", id) } diff --git a/store/store.go b/store/store.go index 6a2a7d6e1d..6a7f386553 100644 --- a/store/store.go +++ b/store/store.go @@ -714,7 +714,7 @@ type FileInfoStore interface { type UploadSessionStore interface { Save(session *model.UploadSession) (*model.UploadSession, error) Update(session *model.UploadSession) error - Get(id string) (*model.UploadSession, error) + Get(ctx context.Context, id string) (*model.UploadSession, error) GetForUser(userID string) ([]*model.UploadSession, error) Delete(id string) error } diff --git a/store/storetest/mocks/UploadSessionStore.go b/store/storetest/mocks/UploadSessionStore.go index dcb8129e9b..f27dbf6856 100644 --- a/store/storetest/mocks/UploadSessionStore.go +++ b/store/storetest/mocks/UploadSessionStore.go @@ -5,6 +5,8 @@ package mocks import ( + context "context" + model "github.com/mattermost/mattermost-server/v6/model" mock "github.com/stretchr/testify/mock" ) @@ -28,13 +30,13 @@ func (_m *UploadSessionStore) Delete(id string) error { return r0 } -// Get provides a mock function with given fields: id -func (_m *UploadSessionStore) Get(id string) (*model.UploadSession, error) { - ret := _m.Called(id) +// Get provides a mock function with given fields: ctx, id +func (_m *UploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { + ret := _m.Called(ctx, id) var r0 *model.UploadSession - if rf, ok := ret.Get(0).(func(string) *model.UploadSession); ok { - r0 = rf(id) + if rf, ok := ret.Get(0).(func(context.Context, string) *model.UploadSession); ok { + r0 = rf(ctx, id) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.UploadSession) @@ -42,8 +44,8 @@ func (_m *UploadSessionStore) Get(id string) (*model.UploadSession, error) { } var r1 error - if rf, ok := ret.Get(1).(func(string) error); ok { - r1 = rf(id) + if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { + r1 = rf(ctx, id) } else { r1 = ret.Error(1) } diff --git a/store/storetest/upload_session_store.go b/store/storetest/upload_session_store.go index ab2178a27f..55d4f89927 100644 --- a/store/storetest/upload_session_store.go +++ b/store/storetest/upload_session_store.go @@ -4,6 +4,7 @@ package storetest import ( + "context" "testing" "time" @@ -52,13 +53,13 @@ func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) { }) t.Run("getting non-existing session should fail", func(t *testing.T) { - us, err := ss.UploadSession().Get("fake") + us, err := ss.UploadSession().Get(context.Background(), "fake") require.Error(t, err) require.Nil(t, us) }) t.Run("getting existing session should succeed", func(t *testing.T) { - us, err := ss.UploadSession().Get(session.Id) + us, err := ss.UploadSession().Get(context.Background(), session.Id) require.NoError(t, err) require.NotNil(t, us) require.Equal(t, session, us) @@ -100,7 +101,7 @@ func testUploadSessionStoreUpdate(t *testing.T, ss store.Store) { err = ss.UploadSession().Update(us) require.NoError(t, err) - updated, err := ss.UploadSession().Get(us.Id) + updated, err := ss.UploadSession().Get(context.Background(), us.Id) require.NoError(t, err) require.NotNil(t, us) require.Equal(t, us, updated) @@ -199,7 +200,7 @@ func testUploadSessionStoreDelete(t *testing.T, ss store.Store) { err = ss.UploadSession().Delete(session.Id) require.NoError(t, err) - us, err = ss.UploadSession().Get(us.Id) + us, err = ss.UploadSession().Get(context.Background(), us.Id) require.Error(t, err) require.Nil(t, us) require.IsType(t, &store.ErrNotFound{}, err) diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index c4c89b935c..55415d1c47 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -9549,10 +9549,10 @@ func (s *TimerLayerUploadSessionStore) Delete(id string) error { return err } -func (s *TimerLayerUploadSessionStore) Get(id string) (*model.UploadSession, error) { +func (s *TimerLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) { start := time.Now() - result, err := s.UploadSessionStore.Get(id) + result, err := s.UploadSessionStore.Get(ctx, id) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil {