From d74825b40e1fc0fafe350dc531d65a7e33e0a223 Mon Sep 17 00:00:00 2001 From: Denis Arh Date: Tue, 13 Nov 2018 13:56:56 +0100 Subject: [PATCH] Add security checks for channel, membership updating --- internal/payload/outgoing.go | 10 ++ internal/payload/outgoing/channel.go | 10 ++ sam/service/channel.go | 223 ++++++++++++--------------- sam/types/channel.go | 10 ++ sam/types/channel_member.go | 20 ++- 5 files changed, 148 insertions(+), 125 deletions(-) diff --git a/internal/payload/outgoing.go b/internal/payload/outgoing.go index 7ca8da32c..06ede7918 100644 --- a/internal/payload/outgoing.go +++ b/internal/payload/outgoing.go @@ -128,6 +128,16 @@ func Channel(ch *samTypes.Channel) *outgoing.Channel { Members: Uint64stoa(ch.Members), View: ChannelView(ch.View), + CanJoin: ch.CanJoin, + CanPart: ch.CanPart, + CanObserve: ch.CanObserve, + CanSendMessages: ch.CanSendMessages, + CanDeleteMessages: ch.CanDeleteMessages, + CanChangeMembers: ch.CanChangeMembers, + CanUpdate: ch.CanUpdate, + CanArchive: ch.CanArchive, + CanDelete: ch.CanDelete, + CreatedAt: ch.CreatedAt, UpdatedAt: ch.UpdatedAt, ArchivedAt: ch.ArchivedAt, diff --git a/internal/payload/outgoing/channel.go b/internal/payload/outgoing/channel.go index c25a78a5f..01e95500f 100644 --- a/internal/payload/outgoing/channel.go +++ b/internal/payload/outgoing/channel.go @@ -39,6 +39,16 @@ type ( Members []string `json:"members,omitempty"` View *ChannelView `json:"view,omitempty"` + CanJoin bool `json:"canJoin"` + CanPart bool `json:"canPart"` + CanObserve bool `json:"canObserve"` + CanSendMessages bool `json:"canSendMessages"` + CanDeleteMessages bool `json:"canDeleteMessages"` + CanChangeMembers bool `json:"canChangeMembers"` + CanUpdate bool `json:"canUpdate"` + CanArchive bool `json:"canArchive"` + CanDelete bool `json:"canDelete"` + CreatedAt time.Time `json:"createdAt"` UpdatedAt *time.Time `json:"updatedAt,omitempty"` ArchivedAt *time.Time `json:"archivedAt,omitempty"` diff --git a/sam/service/channel.go b/sam/service/channel.go index 6013e8266..9c79c8351 100644 --- a/sam/service/channel.go +++ b/sam/service/channel.go @@ -5,6 +5,7 @@ import ( "fmt" "time" + "github.com/davecgh/go-spew/spew" "github.com/pkg/errors" "github.com/crusttech/crust/internal/auth" @@ -38,7 +39,6 @@ type ( Create(channel *types.Channel) (*types.Channel, error) Update(channel *types.Channel) (*types.Channel, error) - FindByMembership() (rval []*types.Channel, err error) FindMembers(channelID uint64) (types.ChannelMemberSet, error) InviteUser(channelID uint64, memberIDs ...uint64) (out types.ChannelMemberSet, err error) @@ -92,13 +92,13 @@ func (svc *channel) FindByID(ID uint64) (ch *types.Channel, err error) { ch, err = svc.channel.FindChannelByID(ID) if err != nil { return + } else if err = svc.preloadExtras(types.ChannelSet{ch}); err != nil { + return } - // @todo [SECURITY] check if user can access this channel - // @todo [SECURITY] check if user can access inactive channels - // if !svc.sec.ch.CanRead(ch) { - // return nil, errors.New("Not allowed to access channel") - // } + if !ch.CanObserve { + return nil, errors.New("Not allowed to access channel") + } return } @@ -113,9 +113,8 @@ func (svc *channel) Find(filter *types.ChannelFilter) (cc types.ChannelSet, err return } - // @todo [SECURITY] check if user can access inactive channels - cc, _ = cc.Filter(func(i *types.Channel) (b bool, e error) { - return true || i.DeletedAt == nil, nil + cc, err = cc.Filter(func(c *types.Channel) (b bool, e error) { + return c.CanObserve, nil }) return @@ -128,6 +127,10 @@ func (svc *channel) preloadExtras(cc types.ChannelSet) (err error) { return } + if err = cc.Walk(svc.setPermissionFlags); err != nil { + return err + } + if err = svc.preloadViews(cc); err != nil { return } @@ -135,19 +138,24 @@ func (svc *channel) preloadExtras(cc types.ChannelSet) (err error) { return } -func (svc *channel) preloadMembers(cc types.ChannelSet) error { - var userID = auth.GetIdentityFromContext(svc.ctx).Identity() +func (svc *channel) preloadMembers(cc types.ChannelSet) (err error) { + var ( + userID = auth.GetIdentityFromContext(svc.ctx).Identity() + mm types.ChannelMemberSet + ) - if mm, err := svc.cmember.Find(&types.ChannelMemberFilter{ComembersOf: userID}); err != nil { - return err + if mm, err = svc.cmember.Find(&types.ChannelMemberFilter{ComembersOf: userID}); err != nil { + return } else { - cc.Walk(func(ch *types.Channel) error { + spew.Dump(mm) + err = cc.Walk(func(ch *types.Channel) error { ch.Members = mm.MembersOf(ch.ID) + ch.Member = mm.FindByChannelID(ch.ID).FindByUserID(userID) return nil }) } - return nil + return } func (svc *channel) preloadViews(cc types.ChannelSet) error { @@ -167,13 +175,11 @@ func (svc *channel) preloadViews(cc types.ChannelSet) error { // FindMembers loads all members (and full users) for a specific channel func (svc *channel) FindMembers(channelID uint64) (out types.ChannelMemberSet, err error) { - var userID = auth.GetIdentityFromContext(svc.ctx).Identity() - - // @todo [SECURITY] check if we can return members on this channel - _ = channelID - _ = userID - return out, svc.db.Transaction(func() (err error) { + if _, err = svc.FindByID(channelID); err != nil { + return + } + out, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}) if err != nil { return err @@ -190,35 +196,6 @@ func (svc *channel) FindMembers(channelID uint64) (out types.ChannelMemberSet, e }) } -// Returns all channels with membership info -func (svc *channel) FindByMembership() (cc []*types.Channel, err error) { - return cc, svc.db.Transaction(func() error { - var chMemberId = repository.Identity(svc.ctx) - - var mm []*types.ChannelMember - - if mm, err = svc.cmember.Find(&types.ChannelMemberFilter{MemberID: chMemberId}); err != nil { - return err - } - - if cc, err = svc.channel.FindChannels(nil); err != nil { - return err - } else if err = svc.preloadExtras(cc); err != nil { - return err - } - - for _, m := range mm { - for _, c := range cc { - if c.ID == m.ChannelID { - c.Member = m - } - } - } - - return nil - }) -} - func (svc *channel) Create(in *types.Channel) (out *types.Channel, err error) { // @todo: [SECURITY] permission check if user can add channel @@ -235,9 +212,11 @@ func (svc *channel) Create(in *types.Channel) (out *types.Channel, err error) { if in.Type == types.ChannelTypeGroup { if out, err = svc.checkGroupExistance(mm); err != nil { return err - } else if out != nil { + } else if out != nil && out.CanObserve { // Group already exists so let's just return it return nil + } else if !out.CanObserve { + return errors.New("Not allowed to create this channel due to permission settings") } } @@ -339,7 +318,7 @@ func (svc *channel) buildMemberSet(owner uint64, members ...uint64) (mm types.Ch // Add all required members and make sure that list is unique for _, m := range members { - if mm.FindByUserId(m) == nil { + if mm.FindByUserID(m) == nil { mm = append(mm, &types.ChannelMember{ UserID: m, Type: types.ChannelMembershipTypeMember, @@ -350,41 +329,40 @@ func (svc *channel) buildMemberSet(owner uint64, members ...uint64) (mm types.Ch return mm } -func (svc *channel) checkGroupExistance(mm types.ChannelMemberSet) (*types.Channel, error) { - if c, err := svc.channel.FindChannelByMemberSet(mm.AllMemberIDs()...); err == repository.ErrChannelNotFound { +func (svc *channel) checkGroupExistance(mm types.ChannelMemberSet) (out *types.Channel, err error) { + if out, err = svc.channel.FindChannelByMemberSet(mm.AllMemberIDs()...); err == repository.ErrChannelNotFound { return nil, nil - } else { - return c, err + } else if out != nil && err == nil { + err = svc.preloadExtras(types.ChannelSet{out}) } + + return } func (svc *channel) Update(in *types.Channel) (ch *types.Channel, err error) { return ch, svc.db.Transaction(func() (err error) { var changed bool - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(in.ID); err != nil { + if ch, err = svc.FindByID(in.ID); err != nil { return } - if ch.ArchivedAt != nil { - return errors.New("Not allowed to edit archived channels") - } else if ch.DeletedAt != nil { - return errors.New("Not allowed to edit deleted channels") + if !ch.CanUpdate { + return errors.New("Not allowed to update this channel") } if in.Type.IsValid() && ch.Type != in.Type { - // @todo [SECURITY] check if user can create public channels + // @todo [SECURITY] check if user can update channel type to public if in.Type == types.ChannelTypePublic && false { return errors.New("Not allowed to change type of this channel to **public**") } - // @todo [SECURITY] check if user can create private channels + // @todo [SECURITY] check if user can create update channel type to private if in.Type == types.ChannelTypePrivate && false { return errors.New("Not allowed to change type of this channel to **private**") } - // @todo [SECURITY] check if user can create group channels + // @todo [SECURITY] check if user can update channel type to group if in.Type == types.ChannelTypeGroup && false { return errors.New("Not allowed to change type of this channel to **group**") } @@ -444,12 +422,13 @@ func (svc *channel) Delete(ID uint64) (ch *types.Channel, err error) { return ch, svc.db.Transaction(func() (err error) { var userID = repository.Identity(svc.ctx) - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(ID); err != nil { + if ch, err = svc.FindByID(ID); err != nil { return } - // @todo [SECURITY] can user delete this channel? + if !ch.CanDelete { + return errors.New("Not allowed to delete this channel") + } if ch.DeletedAt != nil { return errors.New("Channel already deleted") @@ -477,12 +456,13 @@ func (svc *channel) Undelete(ID uint64) (ch *types.Channel, err error) { return ch, svc.db.Transaction(func() (err error) { var userID = repository.Identity(svc.ctx) - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(ID); err != nil { + if ch, err = svc.FindByID(ID); err != nil { return } - // @todo [SECURITY] can user recover this channel? + if !ch.CanDelete { + return errors.New("Not allowed to undelete this channel") + } if ch.DeletedAt == nil { return errors.New("Channel not deleted") @@ -506,12 +486,13 @@ func (svc *channel) Archive(ID uint64) (ch *types.Channel, err error) { return ch, svc.db.Transaction(func() (err error) { var userID = repository.Identity(svc.ctx) - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(ID); err != nil { + if ch, err = svc.FindByID(ID); err != nil { return } - // @todo [SECURITY] can user archive this channel? + if !ch.CanArchive { + return errors.New("Not allowed to archive this channel") + } if ch.ArchivedAt != nil { return errors.New("Channel already archived") @@ -535,12 +516,13 @@ func (svc *channel) Unarchive(ID uint64) (ch *types.Channel, err error) { return ch, svc.db.Transaction(func() (err error) { var userID = repository.Identity(svc.ctx) - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(ID); err != nil { + if ch, err = svc.FindByID(ID); err != nil { return } - // @todo [SECURITY] can user unarchive this channel? + if !ch.CanArchive { + return errors.New("Not allowed to unarchive this channel") + } if ch.ArchivedAt == nil { return errors.New("Channel not archived") @@ -569,8 +551,7 @@ func (svc *channel) InviteUser(channelID uint64, memberIDs ...uint64) (out types out = types.ChannelMemberSet{} - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + if ch, err = svc.FindByID(channelID); err != nil { return } @@ -578,7 +559,9 @@ func (svc *channel) InviteUser(channelID uint64, memberIDs ...uint64) (out types return nil, errors.New("Adding members to a group is not currently supported") } - // @todo [SECURITY] can user add members to this channel? + if !ch.CanChangeMembers { + return nil, errors.New("Not allowed to invite members") + } return out, svc.db.Transaction(func() (err error) { if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { @@ -596,7 +579,7 @@ func (svc *channel) InviteUser(channelID uint64, memberIDs ...uint64) (out types return errors.New("Unexisting user") } - if e := existing.FindByUserId(memberID); e != nil { + if e := existing.FindByUserID(memberID); e != nil { // Already a member/invited e.User = user out = append(out, e) @@ -631,8 +614,7 @@ func (svc *channel) AddMember(channelID uint64, memberIDs ...uint64) (out types. out = types.ChannelMemberSet{} - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + if ch, err = svc.FindByID(channelID); err != nil { return } @@ -640,7 +622,9 @@ func (svc *channel) AddMember(channelID uint64, memberIDs ...uint64) (out types. return nil, errors.New("Adding members to a group is not currently supported") } - // @todo [SECURITY] can user add members to this channel? + if !ch.CanChangeMembers { + return nil, errors.New("Not allowed to add members") + } return out, svc.db.Transaction(func() (err error) { if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { @@ -660,7 +644,7 @@ func (svc *channel) AddMember(channelID uint64, memberIDs ...uint64) (out types. return errors.New("Unexisting user") } - if e := existing.FindByUserId(memberID); e != nil { + if e := existing.FindByUserID(memberID); e != nil { if e.Type != types.ChannelMembershipTypeInvitee { e.User = user out = append(out, e) @@ -716,8 +700,7 @@ func (svc *channel) DeleteMember(channelID uint64, memberIDs ...uint64) (err err existing types.ChannelMemberSet ) - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + if ch, err = svc.FindByID(channelID); err != nil { return } @@ -725,7 +708,9 @@ func (svc *channel) DeleteMember(channelID uint64, memberIDs ...uint64) (err err return errors.New("Removign members from a group is not currently supported") } - // @todo [SECURITY] can user remove members from this channel? + if !ch.CanChangeMembers { + return errors.New("Not allowed to remove members") + } return svc.db.Transaction(func() (err error) { if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { @@ -733,7 +718,7 @@ func (svc *channel) DeleteMember(channelID uint64, memberIDs ...uint64) (err err } for _, memberID := range memberIDs { - if existing.FindByUserId(memberID) == nil { + if existing.FindByUserID(memberID) == nil { // Not really a member... continue } @@ -790,15 +775,20 @@ func (svc *channel) sendChannelEvent(ch *types.Channel) (err error) { // Looks like a valid channel // Preload members, if needed - if len(ch.Members) == 0 { + if len(ch.Members) == 0 || ch.Member == nil { if mm, err := svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: ch.ID}); err != nil { return err } else { ch.Members = mm.AllMemberIDs() + ch.Member = mm.FindByUserID(repository.Identity(svc.ctx)) } } } + if err = svc.setPermissionFlags(ch); err != nil { + return + } + if err = svc.evl.Channel(ch); err != nil { return } @@ -806,36 +796,27 @@ func (svc *channel) sendChannelEvent(ch *types.Channel) (err error) { return nil } -//// @todo temp location, move this somewhere else -//type ( -// nativeChannelSec struct { -// rpo struct { -// ch nativeChannelSecChRepo -// } -// } -// -// nativeChannelSecChRepo interface { -// FindMember(channelId uint64, userId uint64) (*types.User, error) -// } -//) -// -//func ChannelSecurity(chRpo nativeChannelSecChRepo) channelSecurity { -// var sec = &nativeChannelSec{} -// -// sec.rpo.ch = chRpo -// -// return sec -//} -// -//// Current user can read the channel if he is a member -//func (sec nativeChannelSec) CanRead(ch *types.Channel) bool { -// // @todo check if channel is public? -// -// var currentUserID = repository.Identity(svc.ctx) -// -// user, err := sec.rpo.FindMember(ch.ID, currentUserID) -// -// return err != nil && user.Valid() -//} +func (svc *channel) setPermissionFlags(ch *types.Channel) (err error) { + var userID = repository.Identity(svc.ctx) + + var ( + isMember = ch.Member != nil + isCreator = ch.CreatorID == userID + isOwner = isCreator || (isMember && ch.Member.Type == types.ChannelMembershipTypeOwner) + isPublic = ch.Type == types.ChannelTypePublic + ) + + ch.CanJoin = (ch.IsValid() && isPublic) || isOwner + ch.CanPart = isMember + ch.CanObserve = ch.CanJoin + ch.CanSendMessages = ch.CanObserve && isMember + ch.CanDeleteMessages = isOwner + ch.CanChangeMembers = isOwner + ch.CanUpdate = isOwner + ch.CanArchive = isOwner + ch.CanDelete = isOwner + + return nil +} var _ ChannelService = &channel{} diff --git a/sam/types/channel.go b/sam/types/channel.go index a7e0030bd..93abd3c06 100644 --- a/sam/types/channel.go +++ b/sam/types/channel.go @@ -25,6 +25,16 @@ type ( LastMessageID uint64 `json:",omitempty" db:"rel_last_message"` + CanJoin bool `json:"-" db:"-"` + CanPart bool `json:"-" db:"-"` + CanObserve bool `json:"-" db:"-"` + CanSendMessages bool `json:"-" db:"-"` + CanDeleteMessages bool `json:"-" db:"-"` + CanChangeMembers bool `json:"-" db:"-"` + CanUpdate bool `json:"-" db:"-"` + CanArchive bool `json:"-" db:"-"` + CanDelete bool `json:"-" db:"-"` + Member *ChannelMember `json:"-" db:"-"` Members []uint64 `json:"-" db:"-"` View *ChannelView `json:"-" db:"-"` diff --git a/sam/types/channel_member.go b/sam/types/channel_member.go index 2fa22ad3b..edf254db0 100644 --- a/sam/types/channel_member.go +++ b/sam/types/channel_member.go @@ -66,16 +66,28 @@ func (mm ChannelMemberSet) AllMemberIDs() []uint64 { return mmof } -func (uu ChannelMemberSet) FindByUserId(userID uint64) *ChannelMember { - for i := range uu { - if uu[i].UserID == userID { - return uu[i] +func (mm ChannelMemberSet) FindByUserID(userID uint64) *ChannelMember { + for i := range mm { + if mm[i].UserID == userID { + return mm[i] } } return nil } +func (mm ChannelMemberSet) FindByChannelID(channelID uint64) (out ChannelMemberSet) { + out = ChannelMemberSet{} + + for i := range mm { + if mm[i].ChannelID == channelID { + out = append(out, mm[i]) + } + } + + return +} + const ( ChannelMembershipTypeOwner ChannelMembershipType = "owner" ChannelMembershipTypeMember = "member"