diff --git a/sam/db/schema/mysql/20180704080000.base.up.sql b/sam/db/schema/mysql/20180704080000.base.up.sql index 5c8b54c09..d35198bc2 100644 --- a/sam/db/schema/mysql/20180704080000.base.up.sql +++ b/sam/db/schema/mysql/20180704080000.base.up.sql @@ -61,7 +61,7 @@ CREATE TABLE channel_members ( rel_channel BIGINT UNSIGNED NOT NULL REFERENCES channels(id), rel_user BIGINT UNSIGNED NOT NULL, - type ENUM ('owner', 'member') NOT NULL DEFAULT 'member', + type ENUM ('owner', 'member', 'invitee') NOT NULL DEFAULT 'member', created_at DATETIME NOT NULL DEFAULT NOW(), updated_at DATETIME NULL, diff --git a/sam/docs/README.md b/sam/docs/README.md index ea33e6e92..7050f5fc2 100644 --- a/sam/docs/README.md +++ b/sam/docs/README.md @@ -313,7 +313,7 @@ A channel is a representation of a sequence of messages. It has meta data like c | URI | Protocol | Method | Authentication | | --- | -------- | ------ | -------------- | -| `/channels/{channelID}/members/{userID}/join` | HTTP/S | POST | Client ID, Session ID | +| `/channels/{channelID}/members/{userID}` | HTTP/S | PUT | Client ID, Session ID | #### Request parameters @@ -328,7 +328,7 @@ A channel is a representation of a sequence of messages. It has meta data like c | URI | Protocol | Method | Authentication | | --- | -------- | ------ | -------------- | -| `/channels/{channelID}/members/{userID}/part` | HTTP/S | DELETE | Client ID, Session ID | +| `/channels/{channelID}/members/{userID}` | HTTP/S | DELETE | Client ID, Session ID | #### Request parameters diff --git a/sam/docs/src/spec.json b/sam/docs/src/spec.json index aafa7c87d..d7121422e 100644 --- a/sam/docs/src/spec.json +++ b/sam/docs/src/spec.json @@ -270,8 +270,8 @@ }, { "name": "join", - "method": "POST", - "path": "/{channelID}/members/{userID}/join", + "method": "PUT", + "path": "/{channelID}/members/{userID}", "title": "Join channel", "parameters": { "path": [ @@ -283,7 +283,7 @@ { "name": "part", "method": "DELETE", - "path": "/{channelID}/members/{userID}/part", + "path": "/{channelID}/members/{userID}", "title": "Remove member from channel", "parameters": { "path": [ diff --git a/sam/docs/src/spec/channel.json b/sam/docs/src/spec/channel.json index 27b4e6951..2bb7c814d 100644 --- a/sam/docs/src/spec/channel.json +++ b/sam/docs/src/spec/channel.json @@ -142,9 +142,9 @@ }, { "Name": "join", - "Method": "POST", + "Method": "PUT", "Title": "Join channel", - "Path": "/{channelID}/members/{userID}/join", + "Path": "/{channelID}/members/{userID}", "Parameters": { "path": [ { @@ -166,7 +166,7 @@ "Name": "part", "Method": "DELETE", "Title": "Remove member from channel", - "Path": "/{channelID}/members/{userID}/part", + "Path": "/{channelID}/members/{userID}", "Parameters": { "path": [ { diff --git a/sam/repository/channel_member.go b/sam/repository/channel_member.go index 013afae45..393384fc6 100644 --- a/sam/repository/channel_member.go +++ b/sam/repository/channel_member.go @@ -4,7 +4,6 @@ import ( "context" "time" - "github.com/davecgh/go-spew/spew" "github.com/titpetric/factory" "github.com/crusttech/crust/sam/types" @@ -18,7 +17,8 @@ type ( Find(filter *types.ChannelMemberFilter) (types.ChannelMemberSet, error) Create(mod *types.ChannelMember) (*types.ChannelMember, error) - Delete(channelMemberID, userID uint64) error + Update(mod *types.ChannelMember) (*types.ChannelMember, error) + Delete(channelID, userID uint64) error } channelMember struct { @@ -85,8 +85,6 @@ func (r *channelMember) Find(filter *types.ChannelMemberFilter) (types.ChannelMe } } - spew.Dump(filter, sql, params) - return mm, r.db().Select(&mm, sql, params...) } @@ -102,13 +100,13 @@ func (r *channelMember) Create(mod *types.ChannelMember) (*types.ChannelMember, func (r *channelMember) Update(mod *types.ChannelMember) (*types.ChannelMember, error) { mod.UpdatedAt = timeNowPtr() - whitelist := []string{"type", "updated_at"} + whitelist := []string{"type", "updated_at", "rel_channel", "rel_user"} return mod, r.db().UpdatePartial("channel_members", mod, whitelist, "rel_channel", "rel_user") } // Delete removes existing channel membership record -func (r *channelMember) Delete(channelMemberID, userID uint64) error { - sql := `DELETE FROM channel_members WHERE rel_channelMember = ? AND rel_user = ?` - return exec(r.db().Exec(sql, channelMemberID, userID)) +func (r *channelMember) Delete(channelID, userID uint64) error { + sql := `DELETE FROM channel_members WHERE rel_channel = ? AND rel_user = ?` + return exec(r.db().Exec(sql, channelID, userID)) } diff --git a/sam/rest/channel.go b/sam/rest/channel.go index ac9980db3..2fc96b1fc 100644 --- a/sam/rest/channel.go +++ b/sam/rest/channel.go @@ -66,16 +66,16 @@ func (ctrl *Channel) Members(ctx context.Context, r *request.ChannelMembers) (in return ctrl.wrapMemberSet(ctrl.svc.ch.With(ctx).FindMembers(r.ChannelID)) } +func (ctrl *Channel) Invite(ctx context.Context, r *request.ChannelInvite) (interface{}, error) { + return ctrl.wrapMemberSet(ctrl.svc.ch.InviteUser(r.ChannelID, r.UserID...)) +} + func (ctrl *Channel) Join(ctx context.Context, r *request.ChannelJoin) (interface{}, error) { - return nil, nil + return ctrl.wrapMemberSet(ctrl.svc.ch.AddMember(r.ChannelID, r.UserID)) } func (ctrl *Channel) Part(ctx context.Context, r *request.ChannelPart) (interface{}, error) { - return nil, nil -} - -func (ctrl *Channel) Invite(ctx context.Context, r *request.ChannelInvite) (interface{}, error) { - return nil, nil + return nil, ctrl.svc.ch.DeleteMember(r.ChannelID, r.UserID) } func (ctrl *Channel) Attach(ctx context.Context, r *request.ChannelAttach) (interface{}, error) { @@ -117,6 +117,14 @@ func (ctrl *Channel) wrapSet(cc types.ChannelSet, err error) (*outgoing.ChannelS } } +func (ctrl *Channel) wrapMember(m *types.ChannelMember, err error) (*outgoing.ChannelMember, error) { + if err != nil { + return nil, err + } else { + return payload.ChannelMember(m), nil + } +} + func (ctrl *Channel) wrapMemberSet(mm types.ChannelMemberSet, err error) (*outgoing.ChannelMemberSet, error) { if err != nil { return nil, err diff --git a/sam/rest/handlers/channel.go b/sam/rest/handlers/channel.go index aaf57a634..c25a154b6 100644 --- a/sam/rest/handlers/channel.go +++ b/sam/rest/handlers/channel.go @@ -138,8 +138,8 @@ func (ch *Channel) MountRoutes(r chi.Router, middlewares ...func(http.Handler) h r.Get("/{channelID}", ch.Read) r.Delete("/{channelID}", ch.Delete) r.Get("/{channelID}/members", ch.Members) - r.Post("/{channelID}/members/{userID}/join", ch.Join) - r.Delete("/{channelID}/members/{userID}/part", ch.Part) + r.Put("/{channelID}/members/{userID}", ch.Join) + r.Delete("/{channelID}/members/{userID}", ch.Part) r.Post("/{channelID}/invite", ch.Invite) r.Post("/{channelID}/attach", ch.Attach) }) diff --git a/sam/service/channel.go b/sam/service/channel.go index 686899af7..7b91b34a6 100644 --- a/sam/service/channel.go +++ b/sam/service/channel.go @@ -25,6 +25,8 @@ type ( channel repository.ChannelRepository cmember repository.ChannelMemberRepository message repository.MessageRepository + + sysmsgs types.MessageSet } ChannelService interface { @@ -32,12 +34,17 @@ type ( FindByID(channelID uint64) (*types.Channel, error) Find(filter *types.ChannelFilter) (types.ChannelSet, error) - FindByMembership() (rval []*types.Channel, err error) - FindMembers(channelID uint64) (types.ChannelMemberSet, error) 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) + AddMember(channelID uint64, memberIDs ...uint64) (out types.ChannelMemberSet, err error) + DeleteMember(channelID uint64, memberIDs ...uint64) (err error) + Archive(ID uint64) error Unarchive(ID uint64) error Delete(ID uint64) error @@ -67,6 +74,9 @@ func (svc *channel) With(ctx context.Context) ChannelService { channel: repository.Channel(ctx, db), cmember: repository.ChannelMember(ctx, db), message: repository.Message(ctx, db), + + // System messages should be flushed at the end of each session + sysmsgs: types.MessageSet{}, } } @@ -219,12 +229,12 @@ func (svc *channel) Create(in *types.Channel) (out *types.Channel, err error) { // Create the first message, doing this directly with repository to circumvent // message service constraints - msg, err = svc.message.CreateMessage(svc.makeSystemMessage( + svc.scheduleSystemMessage( out, "@%d created new %s channel, topic is: %s", chCreatorID, "", - "")) + "") _ = msg if err != nil { @@ -232,18 +242,16 @@ func (svc *channel) Create(in *types.Channel) (out *types.Channel, err error) { return } - svc.sendMessageEvent(msg) + svc.flushSystemMessages() return svc.evl.Channel(out) }) } -func (svc *channel) Update(in *types.Channel) (out *types.Channel, err error) { +func (svc *channel) Update(ch *types.Channel) (out *types.Channel, err error) { return out, svc.db.Transaction(func() (err error) { - var msgs types.MessageSet - // @todo [SECURITY] can user access this channel? - if out, err = svc.channel.FindChannelByID(in.ID); err != nil { + if out, err = svc.channel.FindChannelByID(ch.ID); err != nil { return } @@ -253,19 +261,19 @@ func (svc *channel) Update(in *types.Channel) (out *types.Channel, err error) { return errors.New("Not allowed to edit deleted channels") } - if out.Type != in.Type { + if out.Type != ch.Type { // @todo [SECURITY] check if user can create public channels - if in.Type == types.ChannelTypePublic && false { + if ch.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 - if in.Type == types.ChannelTypePrivate && false { + if ch.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 - if in.Type == types.ChannelTypeGroup && false { + if ch.Type == types.ChannelTypeGroup && false { return errors.New("Not allowed to change type of this channel to group") } } @@ -273,48 +281,33 @@ func (svc *channel) Update(in *types.Channel) (out *types.Channel, err error) { var chUpdatorId = repository.Identity(svc.ctx) // Copy values - if out.Name != in.Name { + if out.Name != ch.Name { // @todo [SECURITY] can we change channel's name? if false { return errors.New("Not allowed to rename channel") } else { - msgs = append(msgs, svc.makeSystemMessage( - out, "@%d renamed channel %s (was: %s)", chUpdatorId, out.Name, in.Name)) + svc.scheduleSystemMessage(ch, "@%d renamed channel %s (was: %s)", chUpdatorId, out.Name, ch.Name) } - out.Name = in.Name + out.Name = ch.Name } - if out.Topic != in.Topic && true { + if out.Topic != ch.Topic && true { // @todo [SECURITY] can we change channel's topic? if false { return errors.New("Not allowed to change channel topic") } else { - msgs = append(msgs, svc.makeSystemMessage( - out, "@%d changed channel topic: %s (was: %s)", chUpdatorId, out.Topic, in.Topic)) + svc.scheduleSystemMessage(ch, "@%d changed channel topic: %s (was: %s)", chUpdatorId, out.Topic, ch.Topic) } - out.Topic = in.Topic + out.Topic = ch.Topic } // Save the updated channel - if out, err = svc.channel.UpdateChannel(in); err != nil { + if out, err = svc.channel.UpdateChannel(ch); err != nil { return } - // Create the first message, doing this directly with repository to circumvent - // message service constraints - for _, msg := range msgs { - if msg, err = svc.message.CreateMessage(msg); err != nil { - return err - } - - svc.sendMessageEvent(msg) - } - - if err != nil { - // Message creation failed - return - } + svc.flushSystemMessages() return svc.evl.Channel(out) }) @@ -336,16 +329,13 @@ func (svc *channel) Delete(id uint64) error { return errors.New("Channel already deleted") } - msg, err := svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d deleted this channel", userID)) - if err != nil { - return - } - svc.sendMessageEvent(msg) + svc.scheduleSystemMessage(ch, "@%d deleted this channel", userID) if err = svc.channel.DeleteChannelByID(id); err != nil { return } + svc.flushSystemMessages() return svc.evl.Channel(ch) }) } @@ -366,24 +356,20 @@ func (svc *channel) Recover(id uint64) error { return errors.New("Channel not deleted") } - msg, err := svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d recovered this channel", userID)) - if err != nil { - return - } - svc.sendMessageEvent(msg) + svc.scheduleSystemMessage(ch, "@%d recovered this channel", userID) - err = svc.channel.UnarchiveChannelByID(id) - if err != nil { + if err = svc.channel.UnarchiveChannelByID(id); err != nil { return } + svc.flushSystemMessages() return svc.evl.Channel(ch) - }) } func (svc *channel) Archive(id uint64) error { return svc.db.Transaction(func() (err error) { + var userID = repository.Identity(svc.ctx) var ch *types.Channel @@ -398,17 +384,13 @@ func (svc *channel) Archive(id uint64) error { return errors.New("Channel already archived") } - msg, err := svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d archived this channel", userID)) - if err != nil { - return - } - svc.sendMessageEvent(msg) + svc.scheduleSystemMessage(ch, "@%d archived this channel", userID) - err = svc.channel.ArchiveChannelByID(id) - if err != nil { + if err = svc.channel.ArchiveChannelByID(id); err != nil { return } + svc.flushSystemMessages() return svc.evl.Channel(ch) }) } @@ -429,64 +411,208 @@ func (svc *channel) Unarchive(id uint64) error { return errors.New("Channel not archived") } - msg, err := svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d unarchived this channel", userID)) - if err != nil { - return - } - svc.sendMessageEvent(msg) - - err = svc.channel.ArchiveChannelByID(id) - if err != nil { - return - } + svc.scheduleSystemMessage(ch, "@%d unarchived this channel", userID) + svc.flushSystemMessages() return svc.evl.Channel(ch) }) } -func (svc *channel) AddMember(m *types.ChannelMember) (out *types.ChannelMember, err error) { +func (svc *channel) InviteUser(channelID uint64, memberIDs ...uint64) (out types.ChannelMemberSet, err error) { + var ( + userID = repository.Identity(svc.ctx) + ch *types.Channel + existing types.ChannelMemberSet + ) + + out = types.ChannelMemberSet{} + + // @todo [SECURITY] can user access this channel? + if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + return + } + + // @todo [SECURITY] can user add members to this channel? + return out, svc.db.Transaction(func() (err error) { - var userID = repository.Identity(svc.ctx) - - var ch *types.Channel - - // @todo [SECURITY] can user access this channel? - if ch, err = svc.channel.FindChannelByID(m.ChannelID); err != nil { + if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { return } - // @todo [SECURITY] can user add members to this channel? - - msg, err := svc.message.CreateMessage(svc.makeSystemMessage(ch, "@%d added a new member to this channel: @%d", userID, m.UserID)) + users, err := svc.usr.Find(nil) if err != nil { - return + return err } - svc.sendMessageEvent(msg) - return err + for _, memberID := range memberIDs { + user := users.FindById(memberID) + if user == nil { + return errors.New("Unexisting user") + } + + if e := existing.FindByUserId(memberID); e != nil { + // Already a member/invited + e.User = user + out = append(out, e) + continue + } + + svc.scheduleSystemMessage(ch, "@%d invited @%d to the channel", userID, memberID) + + member := &types.ChannelMember{ + ChannelID: channelID, + UserID: memberID, + Type: types.ChannelMembershipTypeInvitee, + } + + if member, err = svc.cmember.Create(member); err != nil { + return err + } + + out = append(out, member) + } + + return svc.flushSystemMessages() }) } -func (svc *channel) makeSystemMessage(ch *types.Channel, format string, a ...interface{}) *types.Message { - return &types.Message{ +func (svc *channel) AddMember(channelID uint64, memberIDs ...uint64) (out types.ChannelMemberSet, err error) { + var ( + userID = repository.Identity(svc.ctx) + ch *types.Channel + existing types.ChannelMemberSet + ) + + out = types.ChannelMemberSet{} + + // @todo [SECURITY] can user access this channel? + if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + return + } + + // @todo [SECURITY] can user add members to this channel? + + return out, svc.db.Transaction(func() (err error) { + if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { + return + } + + users, err := svc.usr.Find(nil) + if err != nil { + return err + } + + for _, memberID := range memberIDs { + var exists bool + + user := users.FindById(memberID) + if user == nil { + return errors.New("Unexisting user") + } + + if e := existing.FindByUserId(memberID); e != nil { + if e.Type != types.ChannelMembershipTypeInvitee { + e.User = user + out = append(out, e) + continue + } else { + exists = true + } + } + + if !exists { + if userID == memberID { + svc.scheduleSystemMessage(ch, "@%d joined", memberID) + } else { + svc.scheduleSystemMessage(ch, "@%d added @%d to the channel", userID, memberID) + } + } + + member := &types.ChannelMember{ + ChannelID: channelID, + UserID: memberID, + Type: types.ChannelMembershipTypeOwner, + } + + if exists { + member, err = svc.cmember.Update(member) + } else { + member, err = svc.cmember.Create(member) + } + + if err != nil { + return err + } + + out = append(out, member) + } + + return svc.flushSystemMessages() + }) +} + +func (svc *channel) DeleteMember(channelID uint64, memberIDs ...uint64) (err error) { + var ( + userID = repository.Identity(svc.ctx) + ch *types.Channel + existing types.ChannelMemberSet + ) + + // @todo [SECURITY] can user access this channel? + if ch, err = svc.channel.FindChannelByID(channelID); err != nil { + return + } + + // @todo [SECURITY] can user remove members from this channel? + + return svc.db.Transaction(func() (err error) { + if existing, err = svc.cmember.Find(&types.ChannelMemberFilter{ChannelID: channelID}); err != nil { + return + } + + for _, memberID := range memberIDs { + if existing.FindByUserId(memberID) == nil { + // Not really a member... + continue + } + + if userID == memberID { + svc.scheduleSystemMessage(ch, "@%d parted", memberID) + } else { + svc.scheduleSystemMessage(ch, "@%d kicked @%d out", userID, memberID) + } + + if err = svc.cmember.Delete(channelID, memberID); err != nil { + return err + } + } + + return svc.flushSystemMessages() + }) +} + +func (svc *channel) scheduleSystemMessage(ch *types.Channel, format string, a ...interface{}) { + svc.sysmsgs = append(svc.sysmsgs, &types.Message{ ChannelID: ch.ID, Message: fmt.Sprintf(format, a...), Type: types.MessageTypeChannelEvent, - } + }) } -// Sends message to event loop -// -// It also preloads user -func (svc *channel) sendMessageEvent(msg *types.Message) (err error) { - if msg.User == nil { - // @todo pull user from cache - if msg.User, err = svc.usr.FindByID(msg.UserID); err != nil { - return +// Flushes sys message stack, stores them into repo & pushes them into event loop +func (svc *channel) flushSystemMessages() (err error) { + defer func() { + svc.sysmsgs = types.MessageSet{} + }() + + return svc.sysmsgs.Walk(func(msg *types.Message) error { + if msg, err = svc.message.CreateMessage(msg); err != nil { + return err + } else { + return svc.evl.Message(msg) } - } + }) - return svc.evl.Message(msg) } //// @todo temp location, move this somewhere else diff --git a/sam/types/channel_member.go b/sam/types/channel_member.go index 6eea3c620..d4e82dba9 100644 --- a/sam/types/channel_member.go +++ b/sam/types/channel_member.go @@ -52,6 +52,16 @@ func (mm ChannelMemberSet) MembersOf(channelID uint64) []uint64 { return mmof } +func (uu ChannelMemberSet) FindByUserId(userID uint64) *ChannelMember { + for i := range uu { + if uu[i].UserID == userID { + return uu[i] + } + } + + return nil +} + const ( ChannelMembershipTypeOwner ChannelMembershipType = "owner" ChannelMembershipTypeMember = "member"