diff --git a/docs/v3-api.yaml b/docs/v3-api.yaml index 11cc9079e..4454f3614 100644 --- a/docs/v3-api.yaml +++ b/docs/v3-api.yaml @@ -714,10 +714,23 @@ paths: Not Found メッセージ、またはスタンプが見つかりません。 operationId: removeMessageStamp + parameters: + - schema: + type: boolean + default: true + in: query + name: include-me + description: "自分が押したスタンプを削除します" + - schema: + type: boolean + default: false + in: query + name: include-other + description: "自分以外が押したスタンプを削除します" tags: - message - stamp - description: 指定したメッセージから指定した自身が押したスタンプを削除します。 + description: 指定したメッセージから指定したスタンプを削除します。 "/stamps/{stampId}": parameters: - $ref: "#/components/parameters/stampIdInPath" @@ -4607,7 +4620,7 @@ paths: post: summary: LiveKit Webhook受信 description: > - LiveKit側で設定したWebhookから呼び出されるエンドポイントです。 + LiveKit側で設定したWebhookから呼び出されるエンドポイントです。 参加者の入室・退出などのイベントを受け取り、サーバ内で処理を行います。 operationId: liveKitWebhook tags: @@ -4630,7 +4643,7 @@ paths: get: summary: サウンドボード用の音声一覧を取得 description: > - DBに保存されたサウンドボード情報を取得します。 + DBに保存されたサウンドボード情報を取得します。 各アイテムには soundId, soundName, stampId が含まれます。 operationId: getSoundboardList tags: @@ -4648,8 +4661,8 @@ paths: post: summary: サウンドボード用の短い音声ファイルをアップロード description: > - 15秒程度の短い音声ファイルを multipart/form-data で送信し、S3(互換ストレージ)にアップロードします。 - クライアントは「soundName」というフィールドを送信し、それをDBに保存して関連付けを行います。 + 15秒程度の短い音声ファイルを multipart/form-data で送信し、S3(互換ストレージ)にアップロードします。 + クライアントは「soundName」というフィールドを送信し、それをDBに保存して関連付けを行います。 また、サーバ側で soundId を自動生成し、S3のファイル名に使用します。 operationId: postSoundboard tags: @@ -4676,8 +4689,8 @@ paths: post: summary: アップロード済み音声を LiveKit ルームで再生 description: > - S3上にある音声ファイルの署名付きURLを生成し、 - Ingressを介して指定ルームに音声を流します。 + S3上にある音声ファイルの署名付きURLを生成し、 + Ingressを介して指定ルームに音声を流します。 該当ルームに参加しているユーザであれば再生可能とします。 operationId: postSoundboardPlay tags: diff --git a/repository/gorm/message.go b/repository/gorm/message.go index d69e9220b..a0c654abb 100644 --- a/repository/gorm/message.go +++ b/repository/gorm/message.go @@ -551,6 +551,42 @@ func (repo *Repository) RemoveStampFromMessage(ctx context.Context, messageID, s return nil } +// RemoveOtherStampFromMessage implements MessageRepository interface. +func (repo *Repository) RemoveOtherStampFromMessage(ctx context.Context, messageID, stampID, userID uuid.UUID) (err error) { + if messageID == uuid.Nil || stampID == uuid.Nil || userID == uuid.Nil { + return repository.ErrNilID + } + + var ms []model.MessageStamp + if err := repo.db.WithContext(ctx).Find(&ms, &model.MessageStamp{MessageID: messageID, StampID: stampID}).Error; err != nil { + return err + } + + result := repo.db.WithContext(ctx). + Where("user_id <> ? AND message_id = ? AND stamp_id = ?", userID, messageID, stampID). + Delete(&model.MessageStamp{}) + if result.Error != nil { + return result.Error + } + + if result.RowsAffected > 0 { + for _, stamp := range ms { + if stamp.UserID == userID { + continue + } + repo.hub.Publish(hub.Message{ + Name: event.MessageUnstamped, + Fields: hub.Fields{ + "message_id": stamp.MessageID, + "stamp_id": stamp.StampID, + "user_id": stamp.UserID, + }, + }) + } + } + return nil +} + func messagePreloads(db *gorm.DB) *gorm.DB { return db. Preload("Stamps"). diff --git a/repository/message.go b/repository/message.go index 913f7eb14..b117873b1 100644 --- a/repository/message.go +++ b/repository/message.go @@ -125,6 +125,12 @@ type MessageRepository interface { // 引数にuuid.Nilを指定するとErrNilIDを返します。 // DBによるエラーを返すことがあります。 RemoveStampFromMessage(ctx context.Context, messageID, stampID, userID uuid.UUID) (err error) + // RemoveOtherStampFromMessage 指定したメッセージから指定したユーザー以外の指定したスタンプを全て削除します + // + // 成功した、或いは既に削除されていた場合、nilを返します。 + // 引数にuuid.Nilを指定するとErrNilIDを返します。 + // DBによるエラーを返すことがあります。 + RemoveOtherStampFromMessage(ctx context.Context, messageID, stampID, userID uuid.UUID) (err error) } // UserUnreadChannel ユーザーの未読チャンネル構造体 diff --git a/repository/mock_repository/mock_message.go b/repository/mock_repository/mock_message.go index 492fb942f..de49b37bf 100644 --- a/repository/mock_repository/mock_message.go +++ b/repository/mock_repository/mock_message.go @@ -204,6 +204,20 @@ func (mr *MockMessageRepositoryMockRecorder) GetUserUnreadChannels(ctx, userID i return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserUnreadChannels", reflect.TypeOf((*MockMessageRepository)(nil).GetUserUnreadChannels), ctx, userID) } +// RemoveOtherStampFromMessage mocks base method. +func (m *MockMessageRepository) RemoveOtherStampFromMessage(ctx context.Context, messageID, stampID, userID uuid.UUID) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "RemoveOtherStampFromMessage", ctx, messageID, stampID, userID) + ret0, _ := ret[0].(error) + return ret0 +} + +// RemoveOtherStampFromMessage indicates an expected call of RemoveOtherStampFromMessage. +func (mr *MockMessageRepositoryMockRecorder) RemoveOtherStampFromMessage(ctx, messageID, stampID, userID interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RemoveOtherStampFromMessage", reflect.TypeOf((*MockMessageRepository)(nil).RemoveOtherStampFromMessage), ctx, messageID, stampID, userID) +} + // RemoveStampFromMessage mocks base method. func (m *MockMessageRepository) RemoveStampFromMessage(ctx context.Context, messageID, stampID, userID uuid.UUID) error { m.ctrl.T.Helper() diff --git a/router/v3/messages.go b/router/v3/messages.go index 9887913c4..f9d342f31 100644 --- a/router/v3/messages.go +++ b/router/v3/messages.go @@ -15,6 +15,11 @@ import ( "github.com/traPtitech/traQ/service/search" ) +type DeleteStampsQuery struct { + IncludeMe string `query:"include-me"` + IncludeOther string `query:"include-other"` +} + // GetMyUnreadChannels GET /users/me/unread func (h *Handlers) GetMyUnreadChannels(c *echo.Context) error { userID := getRequestUserID(c) @@ -285,16 +290,30 @@ func (h *Handlers) AddMessageStamp(c *echo.Context) error { // RemoveMessageStamp DELETE /messages/:messageID/stamps/:stampID func (h *Handlers) RemoveMessageStamp(c *echo.Context) error { + var q DeleteStampsQuery + if err := bindAndValidate(c, &q); err != nil { + return herror.BadRequest(err) + } + + if len(q.IncludeMe) == 0 { + q.IncludeMe = "1" + } + if len(q.IncludeOther) == 0 { + q.IncludeOther = "0" + } + ctx := c.Request().Context() userID := getRequestUserID(c) messageID := getParamAsUUID(c, consts.ParamMessageID) stampID := getParamAsUUID(c, consts.ParamStampID) // スタンプをメッセージから削除 - if err := h.MessageManager.RemoveStamps(ctx, messageID, stampID, userID); err != nil { + if err := h.MessageManager.RemoveStamps(ctx, messageID, stampID, userID, isTrue(q.IncludeMe), isTrue(q.IncludeOther)); err != nil { switch err { case message.ErrChannelArchived: return herror.BadRequest("the channel of this message has been archived") + case message.ErrCannotRemoveStamp: + return herror.Forbidden("you are not allowed to remove this stamp") default: return herror.InternalServerError(err) } diff --git a/service/message/manager.go b/service/message/manager.go index c0b98be4f..92d589c75 100644 --- a/service/message/manager.go +++ b/service/message/manager.go @@ -12,10 +12,11 @@ import ( ) var ( - ErrNotFound = errors.New("not found") - ErrAlreadyExists = errors.New("already exists") - ErrChannelArchived = errors.New("channel archived") - ErrPinLimitExceeded = errors.New("the pin limit exceeded") + ErrNotFound = errors.New("not found") + ErrAlreadyExists = errors.New("already exists") + ErrChannelArchived = errors.New("channel archived") + ErrPinLimitExceeded = errors.New("the pin limit exceeded") + ErrCannotRemoveStamp = errors.New("cannot remove stamps") ) type TimelineQuery struct { @@ -101,13 +102,17 @@ type Manager interface { // 存在しないメッセージを指定した場合は、ErrNotFoundを返します。 // DBによるエラーを返すことがあります。 AddStamps(ctx context.Context, id, stampID, userID uuid.UUID, n int) (*model.MessageStamp, error) - // RemoveStamps 指定したメッセージから指定したユーザーの指定したスタンプを全て削除します + // RemoveStamps 指定したメッセージから指定したスタンプを削除します + // includeMeをtrueに指定すると自分のスタンプを全て削除します。 + // includeOtherをtrueに指定すると自分以外のスタンプを全て削除します。 + // 自分以外のスタンプを削除できるのは、そのメッセージの投稿者であるBotのみです。 // // 成功した場合、或いは既に削除されていた場合、nilを返します。 // アーカイブされているチャンネルを指定すると、ErrChannelArchivedを返します。 // 存在しないメッセージを指定した場合は、ErrNotFoundを返します。 + // スタンプを削除する権限がない場合は、ErrCannotRemoveStampを返します。 // DBによるエラーを返すことがあります。 - RemoveStamps(ctx context.Context, id, stampID, userID uuid.UUID) error + RemoveStamps(ctx context.Context, id, stampID, userID uuid.UUID, includeMe bool, includeOther bool) error Wait(ctx context.Context) error } diff --git a/service/message/manager_impl.go b/service/message/manager_impl.go index 601f7d528..398575f1e 100644 --- a/service/message/manager_impl.go +++ b/service/message/manager_impl.go @@ -293,7 +293,7 @@ func (m *manager) AddStamps(ctx context.Context, id, stampID, userID uuid.UUID, return ms, nil } -func (m *manager) RemoveStamps(ctx context.Context, id, stampID, userID uuid.UUID) error { +func (m *manager) RemoveStamps(ctx context.Context, id, stampID, userID uuid.UUID, includeMe bool, includeOther bool) error { // メッセージ取得 msg, err := m.get(ctx, id) if err != nil { @@ -305,9 +305,27 @@ func (m *manager) RemoveStamps(ctx context.Context, id, stampID, userID uuid.UUI return ErrChannelArchived } + // 自分以外のスタンプを削除できるのは、Botかつ自分のメッセージのみ + if includeOther { + user, err := m.R.GetUser(ctx, userID, false) + if err != nil { + return fmt.Errorf("failed to GetUser: %w", err) + } + if !user.IsBot() || msg.GetUserID() != userID { + return ErrCannotRemoveStamp + } + } + // スタンプを消す - if err := m.R.RemoveStampFromMessage(ctx, id, stampID, userID); err != nil { - return fmt.Errorf("failed to RemoveStampFromMessage: %w", err) + if includeMe { + if err := m.R.RemoveStampFromMessage(ctx, id, stampID, userID); err != nil { + return fmt.Errorf("failed to RemoveStampFromMessage: %w", err) + } + } + if includeOther { + if err := m.R.RemoveOtherStampFromMessage(ctx, id, stampID, userID); err != nil { + return fmt.Errorf("failed to RemoveOtherStampFromMessage: %w", err) + } } // キャッシュ削除 diff --git a/service/message/manager_impl_test.go b/service/message/manager_impl_test.go index fa98d3258..ce9e340f3 100644 --- a/service/message/manager_impl_test.go +++ b/service/message/manager_impl_test.go @@ -469,7 +469,7 @@ func TestManager_RemoveStamps(t *testing.T) { Return(nil, repository.ErrNotFound). Times(1) - err := m.RemoveStamps(context.TODO(), id, uuid.NewV3(uuid.Nil, "s1"), uuid.NewV3(uuid.Nil, "u1")) + err := m.RemoveStamps(context.TODO(), id, uuid.NewV3(uuid.Nil, "s1"), uuid.NewV3(uuid.Nil, "u1"), true, false) assert.EqualError(t, err, ErrNotFound.Error()) }) @@ -488,7 +488,7 @@ func TestManager_RemoveStamps(t *testing.T) { cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) tree.EXPECT().IsArchivedChannel(cid).Return(true).Times(1) - err := m.RemoveStamps(context.TODO(), id, uuid.NewV3(uuid.Nil, "s1"), uuid.NewV3(uuid.Nil, "u1")) + err := m.RemoveStamps(context.TODO(), id, uuid.NewV3(uuid.Nil, "s1"), uuid.NewV3(uuid.Nil, "u1"), true, false) assert.EqualError(t, err, ErrChannelArchived.Error()) }) @@ -540,7 +540,7 @@ func TestManager_RemoveStamps(t *testing.T) { Return(nil). Times(1) - err := m.RemoveStamps(context.TODO(), id, sid, uid) + err := m.RemoveStamps(context.TODO(), id, sid, uid, true, false) if assert.NoError(t, err) { msg, err := m.Get(context.TODO(), id) if assert.NoError(t, err) { @@ -555,4 +555,178 @@ func TestManager_RemoveStamps(t *testing.T) { } } }) + + t.Run("includeOther: non-bot cannot remove others' stamps from another user's message", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + msgOwnerID := uuid.NewV3(uuid.Nil, "u2") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: msgOwnerID}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(&model.User{Bot: false}, nil). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, false, true) + assert.EqualError(t, err, ErrCannotRemoveStamp.Error()) + }) + + t.Run("includeOther: non-bot cannot remove others' stamps even from own message", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: uid}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(&model.User{Bot: false}, nil). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, false, true) + assert.EqualError(t, err, ErrCannotRemoveStamp.Error()) + }) + + t.Run("includeOther: bot cannot remove others' stamps from another user's message", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + msgOwnerID := uuid.NewV3(uuid.Nil, "u2") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: msgOwnerID}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(&model.User{Bot: true}, nil). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, false, true) + assert.EqualError(t, err, ErrCannotRemoveStamp.Error()) + }) + + t.Run("includeOther: bot can remove others' stamps from own message", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: uid}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(&model.User{Bot: true}, nil). + Times(1) + repo.MockMessageRepository. + EXPECT(). + RemoveOtherStampFromMessage(gomock.Any(), id, sid, uid). + Return(nil). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, false, true) + assert.NoError(t, err) + }) + + t.Run("includeOther+includeMe: bot can remove all stamps from own message", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: uid}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(&model.User{Bot: true}, nil). + Times(1) + repo.MockMessageRepository. + EXPECT(). + RemoveStampFromMessage(gomock.Any(), id, sid, uid). + Return(nil). + Times(1) + repo.MockMessageRepository. + EXPECT(). + RemoveOtherStampFromMessage(gomock.Any(), id, sid, uid). + Return(nil). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, true, true) + assert.NoError(t, err) + }) + + t.Run("includeOther: GetUser error is propagated", func(t *testing.T) { + t.Parallel() + ctrl := gomock.NewController(t) + m, cm, repo, tree := setupM(ctrl) + + id := uuid.NewV3(uuid.Nil, "m1") + cid := uuid.NewV3(uuid.Nil, "c1") + sid := uuid.NewV3(uuid.Nil, "s1") + msgOwnerID := uuid.NewV3(uuid.Nil, "u2") + uid := uuid.NewV3(uuid.Nil, "u1") + repo.MockMessageRepository. + EXPECT(). + GetMessageByID(gomock.Any(), id). + Return(&model.Message{ID: id, ChannelID: cid, UserID: msgOwnerID}, nil). + Times(1) + cm.EXPECT().IsPublicChannel(gomock.Any(), cid).Return(true).Times(1) + tree.EXPECT().IsArchivedChannel(cid).Return(false).Times(1) + repo.MockUserRepository. + EXPECT(). + GetUser(gomock.Any(), uid, false). + Return(nil, repository.ErrNotFound). + Times(1) + + err := m.RemoveStamps(context.TODO(), id, sid, uid, false, true) + assert.ErrorContains(t, err, "failed to GetUser") + }) } diff --git a/service/message/manager_test.go b/service/message/manager_test.go index 6ba0abb53..53fb7628e 100644 --- a/service/message/manager_test.go +++ b/service/message/manager_test.go @@ -11,6 +11,7 @@ type Repo struct { *mock_repository.MockChannelRepository *mock_repository.MockMessageRepository *mock_repository.MockPinRepository + *mock_repository.MockUserRepository testutils.EmptyTestRepository } @@ -19,5 +20,6 @@ func NewMockRepo(ctrl *gomock.Controller) *Repo { MockChannelRepository: mock_repository.NewMockChannelRepository(ctrl), MockMessageRepository: mock_repository.NewMockMessageRepository(ctrl), MockPinRepository: mock_repository.NewMockPinRepository(ctrl), + MockUserRepository: mock_repository.NewMockUserRepository(ctrl), } }