diff --git a/api4/drafts.go b/api4/drafts.go index 2dd6210c72..92dd6f2c58 100644 --- a/api4/drafts.go +++ b/api4/drafts.go @@ -120,7 +120,12 @@ func deleteDraft(c *Context, w http.ResponseWriter, r *http.Request) { channelID := c.Params.ChannelId draft, err := c.App.GetDraft(userID, channelID, rootID) - if err != nil || c.AppContext.Session().UserId != draft.UserId { + if err != nil { + c.Err = err + return + } + + if c.AppContext.Session().UserId != draft.UserId { c.SetPermissionError(model.PermissionDeletePost) return } diff --git a/api4/drafts_test.go b/api4/drafts_test.go index 879d78bbb8..9788bd1897 100644 --- a/api4/drafts_test.go +++ b/api4/drafts_test.go @@ -33,7 +33,6 @@ func TestUpsertDraft(t *testing.T) { draft := &model.Draft{ CreateAt: 12345, UpdateAt: 12345, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "original", @@ -105,7 +104,6 @@ func TestGetDrafts(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel1.Id, Message: "draft1", @@ -114,7 +112,6 @@ func TestGetDrafts(t *testing.T) { draft2 := &model.Draft{ CreateAt: 11111, UpdateAt: 32222, - DeleteAt: 0, UserId: user.Id, ChannelId: channel2.Id, Message: "draft2", @@ -180,7 +177,6 @@ func TestDeleteDraft(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel1.Id, Message: "draft1", @@ -190,7 +186,6 @@ func TestDeleteDraft(t *testing.T) { draft2 := &model.Draft{ CreateAt: 11111, UpdateAt: 32222, - DeleteAt: 0, UserId: user.Id, ChannelId: channel2.Id, Message: "draft2", diff --git a/app/draft.go b/app/draft.go index 86786c909f..46f96d7fae 100644 --- a/app/draft.go +++ b/app/draft.go @@ -20,7 +20,7 @@ func (a *App) GetDraft(userID, channelID, rootID string) (*model.Draft, *model.A return nil, model.NewAppError("GetDraft", "app.draft.feature_disabled", nil, "", http.StatusNotImplemented) } - draft, err := a.Srv().Store().Draft().Get(userID, channelID, rootID) + draft, err := a.Srv().Store().Draft().Get(userID, channelID, rootID, false) if err != nil { var nfErr *store.ErrNotFound switch { @@ -39,7 +39,7 @@ func (a *App) UpsertDraft(c *request.Context, draft *model.Draft, connectionID s return nil, model.NewAppError("UpsertDraft", "app.draft.feature_disabled", nil, "", http.StatusNotImplemented) } - dt, dErr := a.Srv().Store().Draft().Get(draft.UserId, draft.ChannelId, draft.RootId) + dt, dErr := a.Srv().Store().Draft().Get(draft.UserId, draft.ChannelId, draft.RootId, true) var notFoundErr *store.ErrNotFound if dErr != nil && !errors.As(dErr, ¬FoundErr) { return nil, model.NewAppError("UpsertDraft", "app.select_error", nil, dErr.Error(), http.StatusInternalServerError) @@ -189,7 +189,7 @@ func (a *App) DeleteDraft(userID, channelID, rootID, connectionID string) (*mode return nil, model.NewAppError("DeleteDraft", "app.draft.feature_disabled", nil, "", http.StatusNotImplemented) } - draft, nErr := a.Srv().Store().Draft().Get(userID, channelID, rootID) + draft, nErr := a.Srv().Store().Draft().Get(userID, channelID, rootID, false) if nErr != nil { return nil, model.NewAppError("DeleteDraft", "app.draft.get.app_error", nil, nErr.Error(), http.StatusBadRequest) } diff --git a/app/draft_test.go b/app/draft_test.go index fc543de443..2ab0983f45 100644 --- a/app/draft_test.go +++ b/app/draft_test.go @@ -34,7 +34,6 @@ func TestGetDraft(t *testing.T) { draft := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft", @@ -84,7 +83,6 @@ func TestUpsertDraft(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft1", @@ -93,7 +91,6 @@ func TestUpsertDraft(t *testing.T) { draft2 := &model.Draft{ CreateAt: 00001, UpdateAt: 00002, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft2", @@ -148,7 +145,6 @@ func TestCreateDraft(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft", @@ -157,7 +153,6 @@ func TestCreateDraft(t *testing.T) { draft2 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel2.Id, Message: "draft2", @@ -223,7 +218,6 @@ func TestUpdateDraft(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft1", @@ -232,7 +226,6 @@ func TestUpdateDraft(t *testing.T) { draft2 := &model.Draft{ CreateAt: 00001, UpdateAt: 00002, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft2", @@ -305,7 +298,6 @@ func TestGetDraftsForUser(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft1", @@ -314,7 +306,6 @@ func TestGetDraftsForUser(t *testing.T) { draft2 := &model.Draft{ CreateAt: 00005, UpdateAt: 00005, - DeleteAt: 0, UserId: user.Id, ChannelId: channel2.Id, Message: "draft2", @@ -400,7 +391,6 @@ func TestDeleteDraft(t *testing.T) { draft1 := &model.Draft{ CreateAt: 00001, UpdateAt: 00001, - DeleteAt: 0, UserId: user.Id, ChannelId: channel.Id, Message: "draft1", diff --git a/app/platform/web_conn.go b/app/platform/web_conn.go index 4e4fc75062..d9fa60cd5d 100644 --- a/app/platform/web_conn.go +++ b/app/platform/web_conn.go @@ -753,15 +753,15 @@ func (wc *WebConn) ShouldSendEvent(msg *model.WebSocketEvent) bool { return wc.GetConnectionID() == msg.GetBroadcast().ConnectionId } + if wc.GetConnectionID() == msg.GetBroadcast().OmitConnectionId { + return false + } + // If the event is destined to a specific user if msg.GetBroadcast().UserId != "" { return wc.UserId == msg.GetBroadcast().UserId } - if wc.GetConnectionID() == msg.GetBroadcast().OmitConnectionId { - return false - } - // if the user is omitted don't send the message if len(msg.GetBroadcast().OmitUsers) > 0 { if _, ok := msg.GetBroadcast().OmitUsers[wc.UserId]; ok { diff --git a/app/web_conn_test.go b/app/web_conn_test.go index 6c05aacf56..7b30d292a7 100644 --- a/app/web_conn_test.go +++ b/app/web_conn_test.go @@ -123,6 +123,7 @@ func TestWebConnShouldSendEvent(t *testing.T) { {"should only send to non-admins", &model.WebsocketBroadcast{ContainsSanitizedData: true}, true, true, false, true}, {"should send to nobody", &model.WebsocketBroadcast{ContainsSensitiveData: true, ContainsSanitizedData: true}, false, false, false, false}, {"should omit basic user 2 by connection id", &model.WebsocketBroadcast{OmitConnectionId: user2ConnID}, true, false, true, true}, + {"should omit basic user 2 by connection id while user is set", &model.WebsocketBroadcast{UserId: th.BasicUser2.Id, OmitConnectionId: user2ConnID}, false, false, false, false}, // needs more cases to get full coverage } diff --git a/db/migrations/migrations.list b/db/migrations/migrations.list index 5ce8502a01..a7fffb38f3 100644 --- a/db/migrations/migrations.list +++ b/db/migrations/migrations.list @@ -198,6 +198,8 @@ db/migrations/mysql/000098_create_post_acknowledgements.down.sql db/migrations/mysql/000098_create_post_acknowledgements.up.sql db/migrations/mysql/000099_create_drafts.down.sql db/migrations/mysql/000099_create_drafts.up.sql +db/migrations/mysql/000100_add_draft_priority_column.down.sql +db/migrations/mysql/000100_add_draft_priority_column.up.sql db/migrations/postgres/000001_create_teams.down.sql db/migrations/postgres/000001_create_teams.up.sql db/migrations/postgres/000002_create_team_members.down.sql @@ -396,3 +398,5 @@ db/migrations/postgres/000098_create_post_acknowledgements.down.sql db/migrations/postgres/000098_create_post_acknowledgements.up.sql db/migrations/postgres/000099_create_drafts.down.sql db/migrations/postgres/000099_create_drafts.up.sql +db/migrations/postgres/000100_add_draft_priority_column.down.sql +db/migrations/postgres/000100_add_draft_priority_column.up.sql diff --git a/db/migrations/mysql/000100_add_draft_priority_column.down.sql b/db/migrations/mysql/000100_add_draft_priority_column.down.sql new file mode 100644 index 0000000000..a4f15cde40 --- /dev/null +++ b/db/migrations/mysql/000100_add_draft_priority_column.down.sql @@ -0,0 +1,14 @@ +SET @preparedStatement = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'Drafts' + AND table_schema = DATABASE() + AND column_name = 'Priority' + ) > 0, + 'ALTER TABLE Drafts DROP COLUMN Priority;', + 'SELECT 1' +)); + +PREPARE alterIfExists FROM @preparedStatement; +EXECUTE alterIfExists; +DEALLOCATE PREPARE alterIfExists; diff --git a/db/migrations/mysql/000100_add_draft_priority_column.up.sql b/db/migrations/mysql/000100_add_draft_priority_column.up.sql new file mode 100644 index 0000000000..134cc86f39 --- /dev/null +++ b/db/migrations/mysql/000100_add_draft_priority_column.up.sql @@ -0,0 +1,14 @@ +SET @preparedStatement = (SELECT IF( + ( + SELECT COUNT(*) FROM INFORMATION_SCHEMA.COLUMNS + WHERE table_name = 'Drafts' + AND table_schema = DATABASE() + AND column_name = 'Priority' + ) > 0, + 'SELECT 1', + 'ALTER TABLE Drafts ADD COLUMN Priority text;' +)); + +PREPARE alterIfExists FROM @preparedStatement; +EXECUTE alterIfExists; +DEALLOCATE PREPARE alterIfExists; diff --git a/db/migrations/postgres/000100_add_draft_priority_column.down.sql b/db/migrations/postgres/000100_add_draft_priority_column.down.sql new file mode 100644 index 0000000000..db071074d2 --- /dev/null +++ b/db/migrations/postgres/000100_add_draft_priority_column.down.sql @@ -0,0 +1 @@ +ALTER TABLE drafts DROP COLUMN IF EXISTS priority; diff --git a/db/migrations/postgres/000100_add_draft_priority_column.up.sql b/db/migrations/postgres/000100_add_draft_priority_column.up.sql new file mode 100644 index 0000000000..4ef6f9e991 --- /dev/null +++ b/db/migrations/postgres/000100_add_draft_priority_column.up.sql @@ -0,0 +1 @@ +ALTER TABLE drafts ADD COLUMN IF NOT EXISTS priority text; diff --git a/i18n/en.json b/i18n/en.json index 365a0d2199..240aa7aac1 100644 --- a/i18n/en.json +++ b/i18n/en.json @@ -8679,6 +8679,10 @@ "id": "model.draft.is_valid.msg.app_error", "translation": "Invalid message." }, + { + "id": "model.draft.is_valid.priority.app_error", + "translation": "Invalid priority" + }, { "id": "model.draft.is_valid.props.app_error", "translation": "Invalid props." diff --git a/model/draft.go b/model/draft.go index e683959b1b..a9741e5727 100644 --- a/model/draft.go +++ b/model/draft.go @@ -23,6 +23,7 @@ type Draft struct { Props StringInterface `json:"props"` // Deprecated: use GetProps() FileIds StringArray `json:"file_ids,omitempty"` Metadata *PostMetadata `json:"metadata,omitempty"` + Priority StringInterface `json:"priority,omitempty"` } func (o *Draft) IsValid(maxDraftSize int) *AppError { @@ -58,6 +59,10 @@ func (o *Draft) IsValid(maxDraftSize int) *AppError { return NewAppError("Drafts.IsValid", "model.draft.is_valid.props.app_error", nil, "channelid="+o.ChannelId, http.StatusBadRequest) } + if utf8.RuneCountInString(StringInterfaceToJSON(o.Priority)) > PostPropsMaxRunes { + return NewAppError("Drafts.IsValid", "model.draft.is_valid.priority.app_error", nil, "channelid="+o.ChannelId, http.StatusBadRequest) + } + return nil } @@ -79,6 +84,7 @@ func (o *Draft) PreSave() { } o.UpdateAt = o.CreateAt + o.DeleteAt = 0 o.PreCommit() } diff --git a/model/websocket_message_test.go b/model/websocket_message_test.go index 14c49fa3ff..e8f9c17ecd 100644 --- a/model/websocket_message_test.go +++ b/model/websocket_message_test.go @@ -215,9 +215,10 @@ func TestWebSocketEventDeepCopy(t *testing.T) { TeamId: "ccc", ContainsSanitizedData: true, ContainsSensitiveData: true, + OmitConnectionId: "ddd", } - ev := NewWebSocketEvent("test", "team", "channel", "user", omitUsers, "") + ev := NewWebSocketEvent("test", "team", "channel", "user", omitUsers, "ddd") ev.Add("post", &Post{}) ev.SetBroadcast(broadcast) diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index 0ce432c3d0..aa41c28ab5 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -3269,7 +3269,7 @@ func (s *OpenTracingLayerDraftStore) Delete(userID string, channelID string, roo return err } -func (s *OpenTracingLayerDraftStore) Get(userID string, channelID string, rootID string) (*model.Draft, error) { +func (s *OpenTracingLayerDraftStore) Get(userID string, channelID string, rootID string, includeDeleted bool) (*model.Draft, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "DraftStore.Get") s.Root.Store.SetContext(newCtx) @@ -3278,7 +3278,7 @@ func (s *OpenTracingLayerDraftStore) Get(userID string, channelID string, rootID }() defer span.Finish() - result, err := s.DraftStore.Get(userID, channelID, rootID) + result, err := s.DraftStore.Get(userID, channelID, rootID, includeDeleted) 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 b05b3dfc21..5aca1535cb 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -3651,11 +3651,11 @@ func (s *RetryLayerDraftStore) Delete(userID string, channelID string, rootID st } -func (s *RetryLayerDraftStore) Get(userID string, channelID string, rootID string) (*model.Draft, error) { +func (s *RetryLayerDraftStore) Get(userID string, channelID string, rootID string, includeDeleted bool) (*model.Draft, error) { tries := 0 for { - result, err := s.DraftStore.Get(userID, channelID, rootID) + result, err := s.DraftStore.Get(userID, channelID, rootID, includeDeleted) if err == nil { return result, nil } diff --git a/store/sqlstore/draft_store.go b/store/sqlstore/draft_store.go index fbf4d07e3d..2dfcc5763d 100644 --- a/store/sqlstore/draft_store.go +++ b/store/sqlstore/draft_store.go @@ -24,7 +24,18 @@ type SqlDraftStore struct { } func draftSliceColumns() []string { - return []string{"CreateAt", "UpdateAt", "DeleteAt", "Message", "RootId", "ChannelId", "UserId", "FileIds", "Props"} + return []string{ + "CreateAt", + "UpdateAt", + "DeleteAt", + "Message", + "RootId", + "ChannelId", + "UserId", + "FileIds", + "Props", + "Priority", + } } func draftToSlice(draft *model.Draft) []interface{} { @@ -38,6 +49,7 @@ func draftToSlice(draft *model.Draft) []interface{} { draft.UserId, model.ArrayToJSON(draft.FileIds), model.StringInterfaceToJSON(draft.Props), + model.StringInterfaceToJSON(draft.Priority), } } @@ -49,17 +61,20 @@ func newSqlDraftStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) } } -func (s *SqlDraftStore) Get(userId, channelId, rootId string) (*model.Draft, error) { +func (s *SqlDraftStore) Get(userId, channelId, rootId string, includeDeleted bool) (*model.Draft, error) { query := s.getQueryBuilder(). - Select("*"). + Select(draftSliceColumns()...). From("Drafts"). Where(sq.Eq{ "UserId": userId, "ChannelId": channelId, "RootId": rootId, - "DeleteAt": 0, }) + if !includeDeleted { + query = query.Where(sq.Eq{"DeleteAt": 0}) + } + dt := model.Draft{} err := s.GetReplicaX().GetBuilder(&dt, query) @@ -108,20 +123,15 @@ func (s *SqlDraftStore) Update(draft *model.Draft) (*model.Draft, error) { Set("Message", draft.Message). Set("Props", draft.Props). Set("FileIds", draft.FileIds). + Set("Priority", draft.Priority). + Set("DeleteAt", 0). Where(sq.Eq{ "UserId": draft.UserId, "ChannelId": draft.ChannelId, "RootId": draft.RootId, - "DeleteAt": 0, }) - sql, args, err := query.ToSql() - - if err != nil { - return nil, errors.Wrapf(err, "failed to convert to sql") - } - - if _, err = s.GetMasterX().Exec(sql, args...); err != nil { + if _, err := s.GetMasterX().ExecBuilder(query); err != nil { return nil, errors.Wrapf(err, "failed to update Draft with channelid=%s", draft.ChannelId) } @@ -132,7 +142,17 @@ func (s *SqlDraftStore) GetDraftsForUser(userID, teamID string) ([]*model.Draft, var drafts []*model.Draft query := s.getQueryBuilder(). - Select("Drafts.*"). + Select( + "Drafts.CreateAt", + "Drafts.UpdateAt", + "Drafts.Message", + "Drafts.RootId", + "Drafts.ChannelId", + "Drafts.UserId", + "Drafts.FileIds", + "Drafts.Props", + "Drafts.Priority", + ). From("Drafts"). InnerJoin("ChannelMembers ON ChannelMembers.ChannelId = Drafts.ChannelId"). Where(sq.And{ diff --git a/store/sqlstore/draft_store_test.go b/store/sqlstore/draft_store_test.go index 471161d0de..25f1a53423 100644 --- a/store/sqlstore/draft_store_test.go +++ b/store/sqlstore/draft_store_test.go @@ -6,345 +6,9 @@ package sqlstore import ( "testing" - "github.com/mattermost/mattermost-server/v6/model" - "github.com/mattermost/mattermost-server/v6/store" "github.com/mattermost/mattermost-server/v6/store/storetest" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" ) func TestDraftStore(t *testing.T) { StoreTestWithSqlStore(t, storetest.TestDraftStore) } - -func TestSaveDraft(t *testing.T) { - StoreTest(t, func(t *testing.T, ss store.Store) { - user := &model.User{ - Id: model.NewId(), - } - - channel := &model.Channel{ - Id: model.NewId(), - } - channel2 := &model.Channel{ - Id: model.NewId(), - } - - member1 := &model.ChannelMember{ - ChannelId: channel.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - member2 := &model.ChannelMember{ - ChannelId: channel2.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - _, err := ss.Channel().SaveMember(member1) - require.NoError(t, err) - - _, err = ss.Channel().SaveMember(member2) - require.NoError(t, err) - - draft1 := &model.Draft{ - CreateAt: 00001, - UpdateAt: 00001, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel.Id, - Message: "draft1", - } - - draft2 := &model.Draft{ - CreateAt: 00005, - UpdateAt: 00005, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel2.Id, - Message: "draft2", - } - - t.Run("save drafts", func(t *testing.T) { - draftResp, err := ss.Draft().Save(draft1) - assert.NoError(t, err) - - assert.Equal(t, draft1.Message, draftResp.Message) - assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) - - draftResp, err = ss.Draft().Save(draft2) - assert.NoError(t, err) - - assert.Equal(t, draft2.Message, draftResp.Message) - assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) - }) - }) -} - -func TestUpdateDraft(t *testing.T) { - StoreTest(t, func(t *testing.T, ss store.Store) { - user := &model.User{ - Id: model.NewId(), - } - - channel := &model.Channel{ - Id: model.NewId(), - } - channel2 := &model.Channel{ - Id: model.NewId(), - } - - member1 := &model.ChannelMember{ - ChannelId: channel.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - member2 := &model.ChannelMember{ - ChannelId: channel2.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - _, err := ss.Channel().SaveMember(member1) - require.NoError(t, err) - - _, err = ss.Channel().SaveMember(member2) - require.NoError(t, err) - - draft1 := &model.Draft{ - CreateAt: 00001, - UpdateAt: 00001, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel.Id, - Message: "draft1", - } - - draft2 := &model.Draft{ - CreateAt: 00005, - UpdateAt: 00005, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel2.Id, - Message: "draft2", - } - - t.Run("update drafts", func(t *testing.T) { - draftResp, err := ss.Draft().Update(draft1) - assert.NoError(t, err) - - assert.Equal(t, draft1.Message, draftResp.Message) - assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) - - draftResp, err = ss.Draft().Update(draft2) - assert.NoError(t, err) - - assert.Equal(t, draft2.Message, draftResp.Message) - assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) - }) - }) -} - -func TestDeleteDraft(t *testing.T) { - StoreTest(t, func(t *testing.T, ss store.Store) { - user := &model.User{ - Id: model.NewId(), - } - - channel := &model.Channel{ - Id: model.NewId(), - } - channel2 := &model.Channel{ - Id: model.NewId(), - } - - member1 := &model.ChannelMember{ - ChannelId: channel.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - member2 := &model.ChannelMember{ - ChannelId: channel2.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - _, err := ss.Channel().SaveMember(member1) - require.NoError(t, err) - - _, err = ss.Channel().SaveMember(member2) - require.NoError(t, err) - - draft1 := &model.Draft{ - CreateAt: 00001, - UpdateAt: 00001, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel.Id, - Message: "draft1", - } - - draft2 := &model.Draft{ - CreateAt: 00005, - UpdateAt: 00005, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel2.Id, - Message: "draft2", - } - - _, err = ss.Draft().Save(draft1) - require.NoError(t, err) - - _, err = ss.Draft().Save(draft2) - require.NoError(t, err) - - t.Run("delete drafts", func(t *testing.T) { - err := ss.Draft().Delete(user.Id, channel.Id, "") - assert.NoError(t, err) - - err = ss.Draft().Delete(user.Id, channel2.Id, "") - assert.NoError(t, err) - }) - }) -} - -func TestGetDraft(t *testing.T) { - StoreTest(t, func(t *testing.T, ss store.Store) { - user := &model.User{ - Id: model.NewId(), - } - - channel := &model.Channel{ - Id: model.NewId(), - } - channel2 := &model.Channel{ - Id: model.NewId(), - } - - member1 := &model.ChannelMember{ - ChannelId: channel.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - member2 := &model.ChannelMember{ - ChannelId: channel2.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - _, err := ss.Channel().SaveMember(member1) - require.NoError(t, err) - - _, err = ss.Channel().SaveMember(member2) - require.NoError(t, err) - - draft1 := &model.Draft{ - CreateAt: 00001, - UpdateAt: 00001, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel.Id, - Message: "draft1", - } - - draft2 := &model.Draft{ - CreateAt: 00005, - UpdateAt: 00005, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel2.Id, - Message: "draft2", - } - - _, err = ss.Draft().Save(draft1) - require.NoError(t, err) - - _, err = ss.Draft().Save(draft2) - require.NoError(t, err) - - t.Run("get drafts", func(t *testing.T) { - draftResp, err := ss.Draft().Get(user.Id, channel.Id, "") - assert.NoError(t, err) - assert.Equal(t, draft1.Message, draftResp.Message) - assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) - - draftResp, err = ss.Draft().Get(user.Id, channel2.Id, "") - assert.NoError(t, err) - assert.Equal(t, draft2.Message, draftResp.Message) - assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) - }) - }) -} - -func TestGetDraftsForUser(t *testing.T) { - StoreTest(t, func(t *testing.T, ss store.Store) { - user := &model.User{ - Id: model.NewId(), - } - - channel := &model.Channel{ - Id: model.NewId(), - } - channel2 := &model.Channel{ - Id: model.NewId(), - } - - member1 := &model.ChannelMember{ - ChannelId: channel.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - member2 := &model.ChannelMember{ - ChannelId: channel2.Id, - UserId: user.Id, - NotifyProps: model.GetDefaultChannelNotifyProps(), - } - - _, err := ss.Channel().SaveMember(member1) - require.NoError(t, err) - - _, err = ss.Channel().SaveMember(member2) - require.NoError(t, err) - - draft1 := &model.Draft{ - CreateAt: 00001, - UpdateAt: 00001, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel.Id, - Message: "draft1", - } - - draft2 := &model.Draft{ - CreateAt: 00005, - UpdateAt: 00005, - DeleteAt: 0, - UserId: user.Id, - ChannelId: channel2.Id, - Message: "draft2", - } - - _, err = ss.Draft().Save(draft1) - require.NoError(t, err) - - _, err = ss.Draft().Save(draft2) - require.NoError(t, err) - - t.Run("get drafts", func(t *testing.T) { - draftResp, err := ss.Draft().GetDraftsForUser(user.Id, "") - assert.NoError(t, err) - - assert.Equal(t, draft2.Message, draftResp[0].Message) - assert.Equal(t, draft2.ChannelId, draftResp[0].ChannelId) - - assert.Equal(t, draft1.Message, draftResp[1].Message) - assert.Equal(t, draft1.ChannelId, draftResp[1].ChannelId) - }) - }) -} diff --git a/store/store.go b/store/store.go index 7cbf0d3e1b..6a2a7d6e1d 100644 --- a/store/store.go +++ b/store/store.go @@ -983,7 +983,7 @@ type PostPriorityStore interface { type DraftStore interface { Save(d *model.Draft) (*model.Draft, error) - Get(userID, channelID, rootID string) (*model.Draft, error) + Get(userID, channelID, rootID string, includeDeleted bool) (*model.Draft, error) Delete(userID, channelID, rootID string) error GetDraftsForUser(userID, teamID string) ([]*model.Draft, error) Update(d *model.Draft) (*model.Draft, error) diff --git a/store/storetest/draft_store.go b/store/storetest/draft_store.go index b540d672d3..9bbe65ee82 100644 --- a/store/storetest/draft_store.go +++ b/store/storetest/draft_store.go @@ -6,8 +6,354 @@ package storetest import ( "testing" + "github.com/mattermost/mattermost-server/v6/model" "github.com/mattermost/mattermost-server/v6/store" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestDraftStore(t *testing.T, ss store.Store, s SqlStore) { + t.Run("SaveDraft", func(t *testing.T) { testSaveDraft(t, ss) }) + t.Run("UpdateDraft", func(t *testing.T) { testUpdateDraft(t, ss) }) + t.Run("DeleteDraft", func(t *testing.T) { testDeleteDraft(t, ss) }) + t.Run("GetDraft", func(t *testing.T) { testGetDraft(t, ss) }) + t.Run("GetDraftsForUser", func(t *testing.T) { testGetDraftsForUser(t, ss) }) +} + +func testSaveDraft(t *testing.T, ss store.Store) { + user := &model.User{ + Id: model.NewId(), + } + + channel := &model.Channel{ + Id: model.NewId(), + } + channel2 := &model.Channel{ + Id: model.NewId(), + } + + member1 := &model.ChannelMember{ + ChannelId: channel.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + member2 := &model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + _, err := ss.Channel().SaveMember(member1) + require.NoError(t, err) + + _, err = ss.Channel().SaveMember(member2) + require.NoError(t, err) + + draft1 := &model.Draft{ + CreateAt: 00001, + UpdateAt: 00001, + UserId: user.Id, + ChannelId: channel.Id, + Message: "draft1", + } + + draft2 := &model.Draft{ + CreateAt: 00005, + UpdateAt: 00005, + UserId: user.Id, + ChannelId: channel2.Id, + Message: "draft2", + } + + t.Run("save drafts", func(t *testing.T) { + draftResp, err := ss.Draft().Save(draft1) + assert.NoError(t, err) + + assert.Equal(t, draft1.Message, draftResp.Message) + assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) + + draftResp, err = ss.Draft().Save(draft2) + assert.NoError(t, err) + + assert.Equal(t, draft2.Message, draftResp.Message) + assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) + }) +} + +func testUpdateDraft(t *testing.T, ss store.Store) { + user := &model.User{ + Id: model.NewId(), + } + + channel := &model.Channel{ + Id: model.NewId(), + } + channel2 := &model.Channel{ + Id: model.NewId(), + } + + member1 := &model.ChannelMember{ + ChannelId: channel.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + member2 := &model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + _, err := ss.Channel().SaveMember(member1) + require.NoError(t, err) + + _, err = ss.Channel().SaveMember(member2) + require.NoError(t, err) + + draft1 := &model.Draft{ + CreateAt: 00001, + UpdateAt: 00001, + UserId: user.Id, + ChannelId: channel.Id, + Message: "draft1", + } + + draft2 := &model.Draft{ + CreateAt: 00005, + UpdateAt: 00005, + UserId: user.Id, + ChannelId: channel2.Id, + Message: "draft2", + } + + t.Run("update drafts", func(t *testing.T) { + draftResp, err := ss.Draft().Update(draft1) + assert.NoError(t, err) + + assert.Equal(t, draft1.Message, draftResp.Message) + assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) + + draftResp, err = ss.Draft().Update(draft2) + assert.NoError(t, err) + + assert.Equal(t, draft2.Message, draftResp.Message) + assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) + }) +} + +func testDeleteDraft(t *testing.T, ss store.Store) { + user := &model.User{ + Id: model.NewId(), + } + + channel := &model.Channel{ + Id: model.NewId(), + } + channel2 := &model.Channel{ + Id: model.NewId(), + } + + member1 := &model.ChannelMember{ + ChannelId: channel.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + member2 := &model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + _, err := ss.Channel().SaveMember(member1) + require.NoError(t, err) + + _, err = ss.Channel().SaveMember(member2) + require.NoError(t, err) + + draft1 := &model.Draft{ + CreateAt: 00001, + UpdateAt: 00001, + UserId: user.Id, + ChannelId: channel.Id, + Message: "draft1", + } + + draft2 := &model.Draft{ + CreateAt: 00005, + UpdateAt: 00005, + UserId: user.Id, + ChannelId: channel2.Id, + Message: "draft2", + } + + _, err = ss.Draft().Save(draft1) + require.NoError(t, err) + + _, err = ss.Draft().Save(draft2) + require.NoError(t, err) + + t.Run("delete drafts", func(t *testing.T) { + err := ss.Draft().Delete(user.Id, channel.Id, "") + assert.NoError(t, err) + + err = ss.Draft().Delete(user.Id, channel2.Id, "") + assert.NoError(t, err) + + _, err = ss.Draft().Get(user.Id, channel.Id, "", false) + require.Error(t, err) + assert.IsType(t, &store.ErrNotFound{}, err) + + _, err = ss.Draft().Get(user.Id, channel2.Id, "", false) + assert.Error(t, err) + assert.IsType(t, &store.ErrNotFound{}, err) + }) +} + +func testGetDraft(t *testing.T, ss store.Store) { + user := &model.User{ + Id: model.NewId(), + } + + channel := &model.Channel{ + Id: model.NewId(), + } + channel2 := &model.Channel{ + Id: model.NewId(), + } + + member1 := &model.ChannelMember{ + ChannelId: channel.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + member2 := &model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + _, err := ss.Channel().SaveMember(member1) + require.NoError(t, err) + + _, err = ss.Channel().SaveMember(member2) + require.NoError(t, err) + + draft1 := &model.Draft{ + CreateAt: 00001, + UpdateAt: 00001, + UserId: user.Id, + ChannelId: channel.Id, + Message: "draft1", + } + + draft2 := &model.Draft{ + CreateAt: 00005, + UpdateAt: 00005, + UserId: user.Id, + ChannelId: channel2.Id, + Message: "draft2", + } + + _, err = ss.Draft().Save(draft1) + require.NoError(t, err) + + _, err = ss.Draft().Save(draft2) + require.NoError(t, err) + + t.Run("get drafts", func(t *testing.T) { + draftResp, err := ss.Draft().Get(user.Id, channel.Id, "", false) + assert.NoError(t, err) + assert.Equal(t, draft1.Message, draftResp.Message) + assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) + + draftResp, err = ss.Draft().Get(user.Id, channel2.Id, "", false) + assert.NoError(t, err) + assert.Equal(t, draft2.Message, draftResp.Message) + assert.Equal(t, draft2.ChannelId, draftResp.ChannelId) + }) + + t.Run("get draft including deleted", func(t *testing.T) { + draftResp, err := ss.Draft().Get(user.Id, channel.Id, "", false) + assert.NoError(t, err) + assert.Equal(t, draft1.Message, draftResp.Message) + assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) + + err = ss.Draft().Delete(user.Id, channel.Id, "") + assert.NoError(t, err) + _, err = ss.Draft().Get(user.Id, channel.Id, "", false) + assert.Error(t, err) + assert.IsType(t, &store.ErrNotFound{}, err) + + draftResp, err = ss.Draft().Get(user.Id, channel.Id, "", true) + assert.NoError(t, err) + assert.Equal(t, draft1.Message, draftResp.Message) + assert.Equal(t, draft1.ChannelId, draftResp.ChannelId) + }) +} + +func testGetDraftsForUser(t *testing.T, ss store.Store) { + user := &model.User{ + Id: model.NewId(), + } + + channel := &model.Channel{ + Id: model.NewId(), + } + channel2 := &model.Channel{ + Id: model.NewId(), + } + + member1 := &model.ChannelMember{ + ChannelId: channel.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + member2 := &model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + + _, err := ss.Channel().SaveMember(member1) + require.NoError(t, err) + + _, err = ss.Channel().SaveMember(member2) + require.NoError(t, err) + + draft1 := &model.Draft{ + CreateAt: 00001, + UpdateAt: 00001, + UserId: user.Id, + ChannelId: channel.Id, + Message: "draft1", + } + + draft2 := &model.Draft{ + CreateAt: 00005, + UpdateAt: 00005, + UserId: user.Id, + ChannelId: channel2.Id, + Message: "draft2", + } + + _, err = ss.Draft().Save(draft1) + require.NoError(t, err) + + _, err = ss.Draft().Save(draft2) + require.NoError(t, err) + + t.Run("get drafts", func(t *testing.T) { + draftResp, err := ss.Draft().GetDraftsForUser(user.Id, "") + assert.NoError(t, err) + + assert.Equal(t, draft2.Message, draftResp[0].Message) + assert.Equal(t, draft2.ChannelId, draftResp[0].ChannelId) + + assert.Equal(t, draft1.Message, draftResp[1].Message) + assert.Equal(t, draft1.ChannelId, draftResp[1].ChannelId) + }) } diff --git a/store/storetest/mocks/DraftStore.go b/store/storetest/mocks/DraftStore.go index b7d89b1308..1eb7d5f17b 100644 --- a/store/storetest/mocks/DraftStore.go +++ b/store/storetest/mocks/DraftStore.go @@ -28,13 +28,13 @@ func (_m *DraftStore) Delete(userID string, channelID string, rootID string) err return r0 } -// Get provides a mock function with given fields: userID, channelID, rootID -func (_m *DraftStore) Get(userID string, channelID string, rootID string) (*model.Draft, error) { - ret := _m.Called(userID, channelID, rootID) +// Get provides a mock function with given fields: userID, channelID, rootID, includeDeleted +func (_m *DraftStore) Get(userID string, channelID string, rootID string, includeDeleted bool) (*model.Draft, error) { + ret := _m.Called(userID, channelID, rootID, includeDeleted) var r0 *model.Draft - if rf, ok := ret.Get(0).(func(string, string, string) *model.Draft); ok { - r0 = rf(userID, channelID, rootID) + if rf, ok := ret.Get(0).(func(string, string, string, bool) *model.Draft); ok { + r0 = rf(userID, channelID, rootID, includeDeleted) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Draft) @@ -42,8 +42,8 @@ func (_m *DraftStore) Get(userID string, channelID string, rootID string) (*mode } var r1 error - if rf, ok := ret.Get(1).(func(string, string, string) error); ok { - r1 = rf(userID, channelID, rootID) + if rf, ok := ret.Get(1).(func(string, string, string, bool) error); ok { + r1 = rf(userID, channelID, rootID, includeDeleted) } else { r1 = ret.Error(1) } diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index da4503ae92..c4c89b935c 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -2996,10 +2996,10 @@ func (s *TimerLayerDraftStore) Delete(userID string, channelID string, rootID st return err } -func (s *TimerLayerDraftStore) Get(userID string, channelID string, rootID string) (*model.Draft, error) { +func (s *TimerLayerDraftStore) Get(userID string, channelID string, rootID string, includeDeleted bool) (*model.Draft, error) { start := time.Now() - result, err := s.DraftStore.Get(userID, channelID, rootID) + result, err := s.DraftStore.Get(userID, channelID, rootID, includeDeleted) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil {