From 56a457be94db3cb52af1951363a8b701aec6818f Mon Sep 17 00:00:00 2001 From: Pugma Date: Mon, 28 Jul 2025 00:52:05 +0900 Subject: [PATCH 1/6] refactor: removed `CheckMessageAccessPerm` from middleware --- router/middlewares/access_control.go | 20 -------------------- router/v3/router.go | 3 +-- 2 files changed, 1 insertion(+), 22 deletions(-) diff --git a/router/middlewares/access_control.go b/router/middlewares/access_control.go index 282b84f54..13b75c44d 100644 --- a/router/middlewares/access_control.go +++ b/router/middlewares/access_control.go @@ -11,7 +11,6 @@ import ( "github.com/traPtitech/traQ/router/extension/herror" "github.com/traPtitech/traQ/service/channel" "github.com/traPtitech/traQ/service/file" - "github.com/traPtitech/traQ/service/message" "github.com/traPtitech/traQ/service/rbac" "github.com/traPtitech/traQ/service/rbac/permission" "github.com/traPtitech/traQ/service/rbac/role" @@ -164,25 +163,6 @@ func CheckClientAccessPerm(rbac rbac.RBAC) echo.MiddlewareFunc { } } -// CheckMessageAccessPerm Messageアクセス権限を確認するミドルウェア -func CheckMessageAccessPerm(cm channel.Manager) echo.MiddlewareFunc { - return func(next echo.HandlerFunc) echo.HandlerFunc { - return func(c echo.Context) error { - userID := c.Get(consts.KeyUser).(model.UserInfo).GetID() - channelID := c.Get(consts.KeyParamMessage).(message.Message).GetChannelID() - - // アクセス権確認 - if ok, err := cm.IsChannelAccessibleToUser(userID, channelID); err != nil { - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() - } - - return next(c) - } - } -} - // CheckChannelAccessPerm Channelアクセス権限を確認するミドルウェア func CheckChannelAccessPerm(cm channel.Manager) echo.MiddlewareFunc { return func(next echo.HandlerFunc) echo.HandlerFunc { diff --git a/router/v3/router.go b/router/v3/router.go index 692b9c697..0905a04ee 100644 --- a/router/v3/router.go +++ b/router/v3/router.go @@ -88,7 +88,6 @@ func (h *Handlers) Setup(e *echo.Group) { requiresWebhookAccessPerm := middlewares.CheckWebhookAccessPerm(h.RBAC) requiresFileAccessPerm := middlewares.CheckFileAccessPerm(h.FileManager) requiresClientAccessPerm := middlewares.CheckClientAccessPerm(h.RBAC) - requiresMessageAccessPerm := middlewares.CheckMessageAccessPerm(h.ChannelManager) requiresChannelAccessPerm := middlewares.CheckChannelAccessPerm(h.ChannelManager) requiresGroupAdminPerm := middlewares.CheckUserGroupAdminPerm(h.RBAC) requiresClipFolderAccessPerm := middlewares.CheckClipFolderAccessPerm() @@ -212,7 +211,7 @@ func (h *Handlers) Setup(e *echo.Group) { apiMessages := api.Group("/messages") { apiMessages.GET("", h.SearchMessages, requires(permission.GetMessage)) - apiMessagesMID := apiMessages.Group("/:messageID", retrieve.MessageID(), requiresMessageAccessPerm) + apiMessagesMID := apiMessages.Group("/:messageID", retrieve.MessageID()) { apiMessagesMID.GET("", h.GetMessage, requires(permission.GetMessage)) apiMessagesMID.PUT("", h.EditMessage, bodyLimit(100), requires(permission.EditMessage)) From 11b6d7b4e84d9b35415eef238cad2619f1f6cd71 Mon Sep 17 00:00:00 2001 From: Pugma Date: Mon, 28 Jul 2025 01:20:49 +0900 Subject: [PATCH 2/6] refactor: impl `IsAccessible` in MessageManager --- service/message/manager.go | 6 ++++++ service/message/manager_impl.go | 10 ++++++++++ 2 files changed, 16 insertions(+) diff --git a/service/message/manager.go b/service/message/manager.go index 30e404043..16fdc4374 100644 --- a/service/message/manager.go +++ b/service/message/manager.go @@ -108,6 +108,12 @@ type Manager interface { // 存在しないメッセージを指定した場合は、ErrNotFoundを返します。 // DBによるエラーを返すことがあります。 RemoveStamps(id, stampID, userID uuid.UUID) error + // IsAccessible 指定したユーザーが指定したメッセージにアクセス可能かどうかを確認します + // + // 成功した場合、アクセス可能かどうかとnilを返します。 + // 存在しないメッセージを指定した場合、falseとErrNotFoundを返します。 + // DBによるエラーを返すことがあります。 + IsAccessible(message Message, userID uuid.UUID) (bool, error) Wait(ctx context.Context) error } diff --git a/service/message/manager_impl.go b/service/message/manager_impl.go index 1e37bbd75..61ab49414 100644 --- a/service/message/manager_impl.go +++ b/service/message/manager_impl.go @@ -316,6 +316,16 @@ func (m *manager) RemoveStamps(id, stampID, userID uuid.UUID) error { return nil } +func (m *manager) IsAccessible(msg Message, userID uuid.UUID) (bool, error) { + // チャンネルアクセス権を確認 + accessible, err := m.CM.IsChannelAccessibleToUser(userID, msg.GetChannelID()) + if err != nil { + return false, fmt.Errorf("failed to check channel access: %w", err) + } + + return accessible, nil +} + func (m *manager) Wait(_ context.Context) error { m.P.Wait() return nil From fe9090801f783b8673033251b48796236c7749c5 Mon Sep 17 00:00:00 2001 From: Pugma Date: Mon, 28 Jul 2025 01:21:59 +0900 Subject: [PATCH 3/6] refactor: use `IsAccessible` in router --- router/v3/messages.go | 132 ++++++++++++++++++++++++++++++++++++++---- 1 file changed, 122 insertions(+), 10 deletions(-) diff --git a/router/v3/messages.go b/router/v3/messages.go index 7c36ff9a1..b3f7e7152 100644 --- a/router/v3/messages.go +++ b/router/v3/messages.go @@ -80,7 +80,20 @@ func (h *Handlers) SearchMessages(c echo.Context) error { // GetMessage GET /messages/:messageID func (h *Handlers) GetMessage(c echo.Context) error { - return c.JSON(http.StatusOK, getParamMessage(c)) + userID := getRequestUserID(c) + m := getParamMessage(c) + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + + return c.JSON(http.StatusOK, m) } // PostMessageRequest POST /channels/:channelID/messages等リクエストボディ @@ -100,6 +113,16 @@ func (h *Handlers) EditMessage(c echo.Context) error { userID := getRequestUserID(c) m := getParamMessage(c) + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + var req PostMessageRequest if err := bindAndValidate(c, &req); err != nil { return err @@ -130,6 +153,16 @@ func (h *Handlers) DeleteMessage(c echo.Context) error { userID := getRequestUserID(c) m := getParamMessage(c) + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + if muid := m.GetUserID(); muid != userID { mUser, err := h.Repo.GetUser(muid, false) if err != nil { @@ -185,7 +218,19 @@ func (h *Handlers) DeleteMessage(c echo.Context) error { // GetPin GET /messages/:messageID/pin func (h *Handlers) GetPin(c echo.Context) error { + userID := getRequestUserID(c) m := getParamMessage(c) + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + if m.GetPin() == nil { return herror.NotFound("this message is not pinned") } @@ -194,8 +239,20 @@ func (h *Handlers) GetPin(c echo.Context) error { // CreatePin POST /messages/:messageID/pin func (h *Handlers) CreatePin(c echo.Context) error { + userID := getRequestUserID(c) m := getParamMessage(c) - p, err := h.MessageManager.Pin(m.GetID(), getRequestUserID(c)) + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + + p, err := h.MessageManager.Pin(m.GetID(), userID) if err != nil { switch err { case message.ErrAlreadyExists: @@ -213,8 +270,20 @@ func (h *Handlers) CreatePin(c echo.Context) error { // RemovePin DELETE /messages/:messageID/pin func (h *Handlers) RemovePin(c echo.Context) error { + userID := getRequestUserID(c) m := getParamMessage(c) - if err := h.MessageManager.Unpin(m.GetID(), getRequestUserID(c)); err != nil { + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + + if err := h.MessageManager.Unpin(m.GetID(), userID); err != nil { switch err { case message.ErrNotFound: return herror.NotFound("pin was not found") @@ -229,7 +298,20 @@ func (h *Handlers) RemovePin(c echo.Context) error { // GetMessageStamps GET /messages/:messageID/stamps func (h *Handlers) GetMessageStamps(c echo.Context) error { - return c.JSON(http.StatusOK, getParamMessage(c).GetStamps()) + userID := getRequestUserID(c) + m := getParamMessage(c) + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + + return c.JSON(http.StatusOK, m.GetStamps()) } // PostMessageStampRequest POST /messages/:messageID/stamps/:stampID リクエストボディ @@ -254,11 +336,21 @@ func (h *Handlers) AddMessageStamp(c echo.Context) error { } userID := getRequestUserID(c) - messageID := getParamAsUUID(c, consts.ParamMessageID) + m := getParamMessage(c) stampID := getParamAsUUID(c, consts.ParamStampID) + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + // スタンプをメッセージに押す - if _, err := h.MessageManager.AddStamps(messageID, stampID, userID, req.Count); err != nil { + if _, err := h.MessageManager.AddStamps(m.GetID(), stampID, userID, req.Count); err != nil { switch err { case message.ErrChannelArchived: return herror.BadRequest("the channel of this message has been archived") @@ -273,11 +365,21 @@ func (h *Handlers) AddMessageStamp(c echo.Context) error { // RemoveMessageStamp DELETE /messages/:messageID/stamps/:stampID func (h *Handlers) RemoveMessageStamp(c echo.Context) error { userID := getRequestUserID(c) - messageID := getParamAsUUID(c, consts.ParamMessageID) + m := getParamMessage(c) stampID := getParamAsUUID(c, consts.ParamStampID) + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + // スタンプをメッセージから削除 - if err := h.MessageManager.RemoveStamps(messageID, stampID, userID); err != nil { + if err := h.MessageManager.RemoveStamps(m.GetID(), stampID, userID); err != nil { switch err { case message.ErrChannelArchived: return herror.BadRequest("the channel of this message has been archived") @@ -292,9 +394,19 @@ func (h *Handlers) RemoveMessageStamp(c echo.Context) error { // GetMessageClips GET /messages/:messageID/clips func (h *Handlers) GetMessageClips(c echo.Context) error { userID := getRequestUserID(c) - messageID := getParamAsUUID(c, consts.ParamMessageID) + m := getParamMessage(c) + + // メッセージアクセス権確認 + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } - clips, err := h.Repo.GetMessageClips(userID, messageID) + clips, err := h.Repo.GetMessageClips(userID, m.GetID()) if err != nil { return herror.InternalServerError(err) } From b4b57cf54d2f66728939bd48d40d02e3058bee32 Mon Sep 17 00:00:00 2001 From: Pugma Date: Mon, 28 Jul 2025 01:43:10 +0900 Subject: [PATCH 4/6] refactor: extracted `checkMessageAccess` in router --- router/v3/messages.go | 104 ++++++++++++++---------------------------- 1 file changed, 34 insertions(+), 70 deletions(-) diff --git a/router/v3/messages.go b/router/v3/messages.go index b3f7e7152..93b2a6a54 100644 --- a/router/v3/messages.go +++ b/router/v3/messages.go @@ -5,6 +5,7 @@ import ( "net/http" vd "github.com/go-ozzo/ozzo-validation/v4" + "github.com/gofrs/uuid" "github.com/labstack/echo/v4" "github.com/traPtitech/traQ/model" @@ -15,6 +16,19 @@ import ( "github.com/traPtitech/traQ/service/search" ) +// checkMessageAccess メッセージアクセス権限を確認するヘルパー関数 +func (h *Handlers) checkMessageAccess(m message.Message, userID uuid.UUID) error { + if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { + if err == message.ErrNotFound { + return herror.NotFound() + } + return herror.InternalServerError(err) + } else if !ok { + return herror.NotFound() + } + return nil +} + // GetMyUnreadChannels GET /users/me/unread func (h *Handlers) GetMyUnreadChannels(c echo.Context) error { userID := getRequestUserID(c) @@ -84,13 +98,8 @@ func (h *Handlers) GetMessage(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } return c.JSON(http.StatusOK, m) @@ -114,13 +123,8 @@ func (h *Handlers) EditMessage(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } var req PostMessageRequest @@ -154,13 +158,8 @@ func (h *Handlers) DeleteMessage(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } if muid := m.GetUserID(); muid != userID { @@ -222,13 +221,8 @@ func (h *Handlers) GetPin(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } if m.GetPin() == nil { @@ -243,13 +237,8 @@ func (h *Handlers) CreatePin(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } p, err := h.MessageManager.Pin(m.GetID(), userID) @@ -274,13 +263,8 @@ func (h *Handlers) RemovePin(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } if err := h.MessageManager.Unpin(m.GetID(), userID); err != nil { @@ -302,13 +286,8 @@ func (h *Handlers) GetMessageStamps(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } return c.JSON(http.StatusOK, m.GetStamps()) @@ -340,13 +319,8 @@ func (h *Handlers) AddMessageStamp(c echo.Context) error { stampID := getParamAsUUID(c, consts.ParamStampID) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } // スタンプをメッセージに押す @@ -369,13 +343,8 @@ func (h *Handlers) RemoveMessageStamp(c echo.Context) error { stampID := getParamAsUUID(c, consts.ParamStampID) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } // スタンプをメッセージから削除 @@ -397,13 +366,8 @@ func (h *Handlers) GetMessageClips(c echo.Context) error { m := getParamMessage(c) // メッセージアクセス権確認 - if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { - return herror.NotFound() - } - return herror.InternalServerError(err) - } else if !ok { - return herror.NotFound() + if err := h.checkMessageAccess(m, userID); err != nil { + return err } clips, err := h.Repo.GetMessageClips(userID, m.GetID()) From 26fc3f9d3cc38b4df224f84c7a08a63caa5f7dde Mon Sep 17 00:00:00 2001 From: Pugma <132571355+Pugma@users.noreply.github.com> Date: Mon, 28 Jul 2025 07:19:58 +0900 Subject: [PATCH 5/6] fix: error handling in `router/v3/messages.go` Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- router/v3/messages.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/router/v3/messages.go b/router/v3/messages.go index 93b2a6a54..d28b4028d 100644 --- a/router/v3/messages.go +++ b/router/v3/messages.go @@ -19,7 +19,7 @@ import ( // checkMessageAccess メッセージアクセス権限を確認するヘルパー関数 func (h *Handlers) checkMessageAccess(m message.Message, userID uuid.UUID) error { if ok, err := h.MessageManager.IsAccessible(m, userID); err != nil { - if err == message.ErrNotFound { + if errors.Is(err, message.ErrNotFound) { return herror.NotFound() } return herror.InternalServerError(err) From c61209d390db757e0ae53f50593d4cca6f328235 Mon Sep 17 00:00:00 2001 From: Pugma Date: Mon, 28 Jul 2025 08:09:03 +0900 Subject: [PATCH 6/6] fix: import required package --- router/v3/messages.go | 1 + 1 file changed, 1 insertion(+) diff --git a/router/v3/messages.go b/router/v3/messages.go index d28b4028d..edfc4158f 100644 --- a/router/v3/messages.go +++ b/router/v3/messages.go @@ -1,6 +1,7 @@ package v3 import ( + "errors" "fmt" "net/http"