Skip to content
Draft
Show file tree
Hide file tree
Changes from 4 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 0 additions & 20 deletions router/middlewares/access_control.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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 {
Expand Down
96 changes: 86 additions & 10 deletions router/v3/messages.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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 {
Comment thread
Pugma marked this conversation as resolved.
Outdated
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)
Expand Down Expand Up @@ -80,7 +94,15 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

return c.JSON(http.StatusOK, m)
}

// PostMessageRequest POST /channels/:channelID/messages等リクエストボディ
Expand All @@ -100,6 +122,11 @@ func (h *Handlers) EditMessage(c echo.Context) error {
userID := getRequestUserID(c)
m := getParamMessage(c)

// メッセージアクセス権確認
if err := h.checkMessageAccess(m, userID); err != nil {
return err
}

var req PostMessageRequest
if err := bindAndValidate(c, &req); err != nil {
return err
Expand Down Expand Up @@ -130,6 +157,11 @@ func (h *Handlers) DeleteMessage(c echo.Context) error {
userID := getRequestUserID(c)
m := getParamMessage(c)

// メッセージアクセス権確認
if err := h.checkMessageAccess(m, userID); err != nil {
return err
}

if muid := m.GetUserID(); muid != userID {
mUser, err := h.Repo.GetUser(muid, false)
if err != nil {
Expand Down Expand Up @@ -185,7 +217,14 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

if m.GetPin() == nil {
return herror.NotFound("this message is not pinned")
}
Expand All @@ -194,8 +233,15 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

p, err := h.MessageManager.Pin(m.GetID(), userID)
if err != nil {
switch err {
case message.ErrAlreadyExists:
Expand All @@ -213,8 +259,15 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

if err := h.MessageManager.Unpin(m.GetID(), userID); err != nil {
switch err {
case message.ErrNotFound:
return herror.NotFound("pin was not found")
Expand All @@ -229,7 +282,15 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

return c.JSON(http.StatusOK, m.GetStamps())
}

// PostMessageStampRequest POST /messages/:messageID/stamps/:stampID リクエストボディ
Expand All @@ -254,11 +315,16 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

// スタンプをメッセージに押す
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")
Expand All @@ -273,11 +339,16 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

// スタンプをメッセージから削除
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")
Expand All @@ -292,9 +363,14 @@ 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 err := h.checkMessageAccess(m, userID); err != nil {
return err
}

clips, err := h.Repo.GetMessageClips(userID, messageID)
clips, err := h.Repo.GetMessageClips(userID, m.GetID())
if err != nil {
return herror.InternalServerError(err)
}
Expand Down
3 changes: 1 addition & 2 deletions router/v3/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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))
Expand Down
6 changes: 6 additions & 0 deletions service/message/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
10 changes: 10 additions & 0 deletions service/message/manager_impl.go
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copilot AI Jul 27, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[nitpick] The function name IsAccessible is ambiguous as it doesn't specify what type of access is being checked. Consider renaming to IsAccessibleToUser or adding more specific documentation about what access permissions are being validated.

Copilot uses AI. Check for mistakes.
}

func (m *manager) Wait(_ context.Context) error {
m.P.Wait()
return nil
Expand Down
Loading